diff --git a/dd-java-agent/instrumentation/java/java-net/java-net-11.0/src/main/java/datadog/trace/instrumentation/httpclient/HttpClientInstrumentation.java b/dd-java-agent/instrumentation/java/java-net/java-net-11.0/src/main/java/datadog/trace/instrumentation/httpclient/HttpClientInstrumentation.java index e30648f61fa..e0a995378de 100644 --- a/dd-java-agent/instrumentation/java/java-net/java-net-11.0/src/main/java/datadog/trace/instrumentation/httpclient/HttpClientInstrumentation.java +++ b/dd-java-agent/instrumentation/java/java-net/java-net-11.0/src/main/java/datadog/trace/instrumentation/httpclient/HttpClientInstrumentation.java @@ -55,6 +55,7 @@ public String[] helperClassNames() { return new String[] { packageName + ".BodyHandlerWrapper", packageName + ".BodyHandlerWrapper$BodySubscriberWrapper", + packageName + ".BodyHandlerWrapper$SubscriptionWrapper", packageName + ".CompletableFutureWrapper", packageName + ".JavaNetClientDecorator", packageName + ".ResponseConsumer" diff --git a/dd-java-agent/instrumentation/java/java-net/java-net-11.0/src/main/java11/datadog/trace/instrumentation/httpclient/BodyHandlerWrapper.java b/dd-java-agent/instrumentation/java/java-net/java-net-11.0/src/main/java11/datadog/trace/instrumentation/httpclient/BodyHandlerWrapper.java index 3a8c7fa70fc..2aa22ccde00 100644 --- a/dd-java-agent/instrumentation/java/java-net/java-net-11.0/src/main/java11/datadog/trace/instrumentation/httpclient/BodyHandlerWrapper.java +++ b/dd-java-agent/instrumentation/java/java-net/java-net-11.0/src/main/java11/datadog/trace/instrumentation/httpclient/BodyHandlerWrapper.java @@ -10,6 +10,7 @@ import java.util.List; import java.util.concurrent.CompletionStage; import java.util.concurrent.Flow; +import java.util.concurrent.atomic.AtomicReferenceFieldUpdater; public class BodyHandlerWrapper implements BodyHandler { private final BodyHandler delegate; @@ -27,12 +28,17 @@ public BodySubscriber apply(ResponseInfo responseInfo) { if (subscriber instanceof BodySubscriberWrapper) { return subscriber; } - return new BodySubscriberWrapper<>(subscriber, span.captureWithContext()); + return new BodySubscriberWrapper<>(subscriber, span.captureWithContext().hold()); } static class BodySubscriberWrapper implements BodySubscriber { + private static final AtomicReferenceFieldUpdater + CONTINUATION = + AtomicReferenceFieldUpdater.newUpdater( + BodySubscriberWrapper.class, ContextContinuation.class, "continuation"); + private final BodySubscriber delegate; - private final ContextContinuation continuation; + private volatile ContextContinuation continuation; public BodySubscriberWrapper(BodySubscriber delegate, ContextContinuation continuation) { this.delegate = delegate; @@ -50,27 +56,87 @@ public CompletionStage getBody() { @Override public void onSubscribe(Flow.Subscription subscription) { - delegate.onSubscribe(subscription); + boolean completed = false; + try { + delegate.onSubscribe(new SubscriptionWrapper(subscription, this)); + completed = true; + } finally { + if (!completed) { + releaseContinuation(); + } + } } @Override public void onNext(List item) { - try (ContextScope ignore = continuation.resume()) { - delegate.onNext(item); + boolean completed = false; + try { + try (ContextScope ignore = resumeContinuation()) { + delegate.onNext(item); + } + completed = true; + } finally { + if (!completed) { + releaseContinuation(); + } } } @Override public void onError(Throwable throwable) { - try (ContextScope ignore = continuation.resume()) { - delegate.onError(throwable); + try { + try (ContextScope ignore = resumeContinuation()) { + delegate.onError(throwable); + } + } finally { + releaseContinuation(); } } @Override public void onComplete() { - try (ContextScope ignore = continuation.resume()) { - delegate.onComplete(); + try { + try (ContextScope ignore = resumeContinuation()) { + delegate.onComplete(); + } + } finally { + releaseContinuation(); + } + } + + private ContextScope resumeContinuation() { + ContextContinuation continuation = this.continuation; + return continuation == null ? null : continuation.resume(); + } + + private void releaseContinuation() { + ContextContinuation continuation = CONTINUATION.getAndSet(this, null); + if (continuation != null) { + continuation.release(); + } + } + } + + static final class SubscriptionWrapper implements Flow.Subscription { + private final Flow.Subscription delegate; + private final BodySubscriberWrapper subscriber; + + SubscriptionWrapper(Flow.Subscription delegate, BodySubscriberWrapper subscriber) { + this.delegate = delegate; + this.subscriber = subscriber; + } + + @Override + public void request(long count) { + delegate.request(count); + } + + @Override + public void cancel() { + try { + delegate.cancel(); + } finally { + subscriber.releaseContinuation(); } } } diff --git a/dd-java-agent/instrumentation/java/java-net/java-net-11.0/src/test/java/datadog/trace/instrumentation/httpclient/BodyHandlerWrapperTest.java b/dd-java-agent/instrumentation/java/java-net/java-net-11.0/src/test/java/datadog/trace/instrumentation/httpclient/BodyHandlerWrapperTest.java new file mode 100644 index 00000000000..ddb930c85f6 --- /dev/null +++ b/dd-java-agent/instrumentation/java/java-net/java-net-11.0/src/test/java/datadog/trace/instrumentation/httpclient/BodyHandlerWrapperTest.java @@ -0,0 +1,190 @@ +package datadog.trace.instrumentation.httpclient; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertSame; +import static org.junit.jupiter.api.Assertions.assertThrows; + +import datadog.context.Context; +import datadog.context.ContextContinuation; +import datadog.context.ContextKey; +import datadog.context.ContextScope; +import datadog.trace.bootstrap.instrumentation.api.AgentSpan; +import java.lang.reflect.Proxy; +import java.net.http.HttpResponse.BodySubscriber; +import java.nio.ByteBuffer; +import java.util.ArrayList; +import java.util.List; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.CompletionStage; +import java.util.concurrent.Flow; +import org.junit.jupiter.api.Test; + +class BodyHandlerWrapperTest { + + @Test + void holdsContextAcrossCallbacksUntilCompletion() { + Context original = Context.current(); + Context captured = original.with(ContextKey.named("response-body"), new Object()); + ContextContinuation continuation = captured.capture(); + RecordingSubscriber subscriber = new RecordingSubscriber(); + BodySubscriber wrapper = wrap(subscriber, continuation); + + try { + wrapper.onNext(List.of()); + assertSame(captured, subscriber.callbackContexts.get(0)); + assertSame(original, Context.current()); + + wrapper.onNext(List.of()); + assertSame(captured, subscriber.callbackContexts.get(1)); + assertSame(original, Context.current()); + + wrapper.onComplete(); + assertSame(captured, subscriber.callbackContexts.get(2)); + assertSame(original, Context.current()); + + // A released continuation must no longer reactivate the captured context. + try (ContextScope ignored = continuation.resume()) { + assertSame(original, Context.current()); + } + } finally { + continuation.release(); + } + } + + @Test + void releasesContinuationWhenOnSubscribeThrows() { + RecordingContinuation continuation = new RecordingContinuation(); + RecordingSubscriber subscriber = new RecordingSubscriber(); + subscriber.throwOnSubscribe = true; + BodySubscriber wrapper = wrap(subscriber, continuation); + + assertThrows( + IllegalStateException.class, () -> wrapper.onSubscribe(new RecordingSubscription())); + assertEquals(1, continuation.released); + wrapper.onComplete(); + + assertEquals(1, continuation.released); + } + + @Test + void releasesContinuationWhenOnNextThrows() { + RecordingContinuation continuation = new RecordingContinuation(); + RecordingSubscriber subscriber = new RecordingSubscriber(); + subscriber.throwOnNext = true; + BodySubscriber wrapper = wrap(subscriber, continuation); + + assertThrows(IllegalStateException.class, () -> wrapper.onNext(List.of())); + assertEquals(1, continuation.released); + wrapper.onComplete(); + + assertEquals(1, continuation.released); + } + + @Test + void releasesContinuationWhenSubscriptionIsCancelled() { + RecordingContinuation continuation = new RecordingContinuation(); + RecordingSubscriber subscriber = new RecordingSubscriber(); + BodySubscriber wrapper = wrap(subscriber, continuation); + RecordingSubscription subscription = new RecordingSubscription(); + + wrapper.onSubscribe(subscription); + subscriber.subscription.request(3); + subscriber.subscription.cancel(); + wrapper.onComplete(); + + assertEquals(3, subscription.requested); + assertEquals(1, subscription.cancelled); + assertEquals(1, continuation.released); + } + + private static BodySubscriber wrap( + RecordingSubscriber subscriber, ContextContinuation continuation) { + AgentSpan span = + (AgentSpan) + Proxy.newProxyInstance( + AgentSpan.class.getClassLoader(), + new Class[] {AgentSpan.class}, + (proxy, method, args) -> continuation); + return new BodyHandlerWrapper<>(ignored -> subscriber, span).apply(null); + } + + private static final class RecordingSubscriber + implements java.net.http.HttpResponse.BodySubscriber { + private final CompletableFuture body = new CompletableFuture<>(); + private final List callbackContexts = new ArrayList<>(); + private Flow.Subscription subscription; + private boolean throwOnSubscribe; + private boolean throwOnNext; + + @Override + public CompletionStage getBody() { + return body; + } + + @Override + public void onSubscribe(Flow.Subscription subscription) { + this.subscription = subscription; + if (throwOnSubscribe) { + throw new IllegalStateException("onSubscribe"); + } + } + + @Override + public void onNext(List item) { + callbackContexts.add(Context.current()); + if (throwOnNext) { + throw new IllegalStateException("onNext"); + } + } + + @Override + public void onError(Throwable throwable) { + body.completeExceptionally(throwable); + } + + @Override + public void onComplete() { + callbackContexts.add(Context.current()); + body.complete(null); + } + } + + private static final class RecordingSubscription implements Flow.Subscription { + private long requested; + private int cancelled; + + @Override + public void request(long count) { + requested += count; + } + + @Override + public void cancel() { + cancelled++; + } + } + + private static final class RecordingContinuation implements ContextContinuation { + private int released; + + @Override + public ContextContinuation hold() { + return this; + } + + @Override + public Context context() { + return null; + } + + @Override + public ContextScope resume() { + return null; + } + + @Override + public void release() { + released++; + } + } +}