diff --git a/dd-java-agent/instrumentation/netty/netty-4.0/src/main/java/datadog/trace/instrumentation/netty40/server/MaybeBlockResponseHandler.java b/dd-java-agent/instrumentation/netty/netty-4.0/src/main/java/datadog/trace/instrumentation/netty40/server/MaybeBlockResponseHandler.java index d7ec1c109ef..6f6c1891e88 100644 --- a/dd-java-agent/instrumentation/netty/netty-4.0/src/main/java/datadog/trace/instrumentation/netty40/server/MaybeBlockResponseHandler.java +++ b/dd-java-agent/instrumentation/netty/netty-4.0/src/main/java/datadog/trace/instrumentation/netty40/server/MaybeBlockResponseHandler.java @@ -70,7 +70,7 @@ public void write(ChannelHandlerContext ctx, Object msg, ChannelPromise prm) thr if (isAnalyzedResponse(channel)) { if (isBlockedResponse(channel)) { // block further writes - log.debug("Write suppressed, msg {} dropped", msg); + log.debug("Write suppressed; dropped outbound message"); ReferenceCountUtil.release(msg); } else { super.write(ctx, msg, prm); diff --git a/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/ChannelFutureListenerInstrumentation.java b/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/ChannelFutureListenerInstrumentation.java index 8fc6c413406..e5ae3ded856 100644 --- a/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/ChannelFutureListenerInstrumentation.java +++ b/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/ChannelFutureListenerInstrumentation.java @@ -46,6 +46,7 @@ public ElementMatcher hierarchyMatcher() { public String[] helperClassNames() { return new String[] { packageName + ".AttributeKeys", + packageName + ".ServerRequestContext", // client helpers packageName + ".client.NettyHttpClientDecorator", packageName + ".client.NettyResponseInjectAdapter", @@ -58,6 +59,7 @@ public String[] helperClassNames() { packageName + ".server.NettyHttpServerDecorator$NettyBlockResponseFunction", packageName + ".server.BlockingResponseHandler", packageName + ".server.BlockingResponseHandler$IgnoreAllWritesHandler", + packageName + ".server.BlockingResponseHandler$PendingBlockResponse", packageName + ".server.HttpServerRequestTracingHandler", packageName + ".server.HttpServerResponseTracingHandler", packageName + ".server.HttpServerTracingHandler" diff --git a/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/NettyChannelHandlerContextInstrumentation.java b/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/NettyChannelHandlerContextInstrumentation.java index da93dff0081..7af9cfb8f33 100644 --- a/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/NettyChannelHandlerContextInstrumentation.java +++ b/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/NettyChannelHandlerContextInstrumentation.java @@ -13,7 +13,6 @@ import static net.bytebuddy.matcher.ElementMatchers.isPublic; import com.google.auto.service.AutoService; -import datadog.context.Context; import datadog.trace.agent.tooling.Instrumenter; import datadog.trace.agent.tooling.InstrumenterModule; import datadog.trace.bootstrap.instrumentation.api.AgentScope; @@ -47,12 +46,14 @@ public ElementMatcher hierarchyMatcher() { public String[] helperClassNames() { return new String[] { packageName + ".AttributeKeys", + packageName + ".ServerRequestContext", packageName + ".client.NettyHttpClientDecorator", packageName + ".server.ResponseExtractAdapter", packageName + ".server.NettyHttpServerDecorator", packageName + ".server.NettyHttpServerDecorator$NettyBlockResponseFunction", packageName + ".server.BlockingResponseHandler", packageName + ".server.BlockingResponseHandler$IgnoreAllWritesHandler", + packageName + ".server.BlockingResponseHandler$PendingBlockResponse", packageName + ".server.HttpServerRequestTracingHandler", packageName + ".server.HttpServerResponseTracingHandler", packageName + ".server.HttpServerTracingHandler" @@ -70,8 +71,8 @@ public void methodAdvice(MethodTransformer transformer) { public static class FireAdvice { @Advice.OnMethodEnter(suppress = Throwable.class) public static AgentScope scopeSpan(@Advice.This final ChannelHandlerContext ctx) { - final Context storedContext = ctx.channel().attr(CONTEXT_ATTRIBUTE_KEY).get(); - final AgentSpan channelSpan = spanFromContext(storedContext); + final AgentSpan channelSpan = + spanFromContext(ctx.channel().attr(CONTEXT_ATTRIBUTE_KEY).get()); if (channelSpan == null || channelSpan == activeSpan()) { // don't modify the scope return null; diff --git a/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/NettyChannelPipelineInstrumentation.java b/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/NettyChannelPipelineInstrumentation.java index a1ceb822132..6fc2bc06f10 100644 --- a/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/NettyChannelPipelineInstrumentation.java +++ b/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/NettyChannelPipelineInstrumentation.java @@ -72,6 +72,7 @@ public ElementMatcher hierarchyMatcher() { public String[] helperClassNames() { return new String[] { packageName + ".AttributeKeys", + packageName + ".ServerRequestContext", // client helpers packageName + ".client.NettyHttpClientDecorator", packageName + ".client.NettyResponseInjectAdapter", @@ -84,6 +85,7 @@ public String[] helperClassNames() { packageName + ".server.NettyHttpServerDecorator$NettyBlockResponseFunction", packageName + ".server.BlockingResponseHandler", packageName + ".server.BlockingResponseHandler$IgnoreAllWritesHandler", + packageName + ".server.BlockingResponseHandler$PendingBlockResponse", packageName + ".server.HttpServerContextTrackingHandler", packageName + ".server.HttpServerRequestTracingHandler", packageName + ".server.HttpServerResponseTracingHandler", diff --git a/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/server/BlockingResponseHandler.java b/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/server/BlockingResponseHandler.java index 1c49fceae2f..3903380a810 100644 --- a/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/server/BlockingResponseHandler.java +++ b/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/server/BlockingResponseHandler.java @@ -4,6 +4,7 @@ import datadog.trace.api.gateway.Flow; import datadog.trace.api.internal.TraceSegment; import datadog.trace.bootstrap.blocking.BlockingActionHelper; +import datadog.trace.instrumentation.netty41.ServerRequestContext; import io.netty.channel.ChannelHandler; import io.netty.channel.ChannelHandlerContext; import io.netty.channel.ChannelInboundHandlerAdapter; @@ -15,6 +16,7 @@ import io.netty.handler.codec.http.HttpRequest; import io.netty.handler.codec.http.HttpResponseStatus; import io.netty.handler.codec.http.HttpUtil; +import io.netty.handler.codec.http.HttpVersion; import io.netty.util.ReferenceCountUtil; import java.util.Map; import org.slf4j.Logger; @@ -22,6 +24,9 @@ public class BlockingResponseHandler extends ChannelInboundHandlerAdapter { public static final Logger log = LoggerFactory.getLogger(BlockingResponseHandler.class); + private static final String IGNORE_ALL_WRITES_HANDLER = "ignore_all_writes_handler"; + private static final String MISSING_RESPONSE_TRACING_HANDLER_MESSAGE = + "Unable to block because HttpServerResponseTracingHandler was not found on the pipeline"; private static volatile boolean HAS_WARNED; private final TraceSegment segment; @@ -29,6 +34,7 @@ public class BlockingResponseHandler extends ChannelInboundHandlerAdapter { private final BlockingContentType bct; private final Map extraHeaders; private final String securityResponseId; + private final ServerRequestContext serverContext; private boolean hasBlockedAlready; @@ -37,21 +43,26 @@ public BlockingResponseHandler( int statusCode, BlockingContentType bct, Map extraHeaders, - String securityResponseId) { + String securityResponseId, + ServerRequestContext serverContext) { this.segment = segment; this.statusCode = statusCode; this.bct = bct; this.extraHeaders = extraHeaders; this.securityResponseId = securityResponseId; + this.serverContext = serverContext; } - public BlockingResponseHandler(TraceSegment segment, Flow.Action.RequestBlockingAction rba) { - this( - segment, - rba.getStatusCode(), - rba.getBlockingContentType(), - rba.getExtraHeaders(), - rba.getSecurityResponseId()); + public BlockingResponseHandler( + TraceSegment segment, + Flow.Action.RequestBlockingAction rba, + ServerRequestContext serverContext) { + this.segment = segment; + this.statusCode = rba.getStatusCode(); + this.bct = rba.getBlockingContentType(); + this.extraHeaders = rba.getExtraHeaders(); + this.securityResponseId = rba.getSecurityResponseId(); + this.serverContext = serverContext; } @Override @@ -73,75 +84,141 @@ public void channelRead(final ChannelHandlerContext ctx, final Object msg) { } if (ctxForDownstream == null) { - if (HAS_WARNED) { - log.debug( - "Unable to block because HttpServerResponseTracingHandler was not found on the pipeline"); - } else { - log.warn( - "Unable to block because HttpServerResponseTracingHandler was not found on the pipeline"); - HAS_WARNED = true; - } + logMissingResponseTracingHandler(); ctx.fireChannelRead(msg); return; } HttpRequest request = (HttpRequest) msg; - int httpCode = BlockingActionHelper.getHttpCode(statusCode); - HttpResponseStatus httpResponseStatus = HttpResponseStatus.valueOf(httpCode); - FullHttpResponse response = - new DefaultFullHttpResponse(request.protocolVersion(), httpResponseStatus); + this.hasBlockedAlready = true; - HttpHeaders headers = response.headers(); - headers.set("Connection", "close"); + PendingBlockResponse pendingBlockResponse = + new PendingBlockResponse( + segment, + statusCode, + bct, + extraHeaders, + securityResponseId, + request.protocolVersion(), + request.headers().get("accept")); + ReferenceCountUtil.release(msg); - for (Map.Entry h : this.extraHeaders.entrySet()) { - headers.set(h.getKey(), h.getValue()); + if (serverContext != null + && ServerRequestContext.nextResponse(ctx.channel()) != serverContext) { + serverContext.deferBlockResponse(pendingBlockResponse); + return; } - if (bct != BlockingContentType.NONE) { - String acceptHeader = request.headers().get("accept"); - BlockingActionHelper.TemplateType type = - BlockingActionHelper.determineTemplateType(bct, acceptHeader); - headers.set("Content-type", BlockingActionHelper.getContentType(type)); + writeBlockResponse(ctxForDownstream, pendingBlockResponse); + } - byte[] template = BlockingActionHelper.getTemplate(type, this.securityResponseId); - HttpUtil.setContentLength(response, template.length); - response.content().writeBytes(template); + private static void logMissingResponseTracingHandler() { + if (HAS_WARNED) { + log.debug(MISSING_RESPONSE_TRACING_HANDLER_MESSAGE); + } else { + log.warn(MISSING_RESPONSE_TRACING_HANDLER_MESSAGE); + HAS_WARNED = true; } + } - this.hasBlockedAlready = true; - - ReferenceCountUtil.release(msg); + static boolean maybeWriteDeferredBlockResponse( + ChannelHandlerContext ctx, ServerRequestContext serverContext) { + if (serverContext == null) { + return false; + } + Object deferredBlockResponse = serverContext.deferredBlockResponse(); + if (!(deferredBlockResponse instanceof PendingBlockResponse)) { + return false; + } + serverContext.deferBlockResponse(null); + writeBlockResponse(ctx, (PendingBlockResponse) deferredBlockResponse); + return true; + } + private static void writeBlockResponse( + ChannelHandlerContext ctxForDownstream, PendingBlockResponse pendingBlockResponse) { // write starts in the handler before the one associated with ctx // so add one that will be skipped (but that will prevent any writes later coming from later // handlers). // We do not want to start from the end of the // pipeline because there is an increased risk of hitting duplex handlers that // expect to have seen a request before processing the response - ctxForDownstream = - ctxForDownstream - .pipeline() - .addAfter( - ctxForDownstream.name(), - "ignore_all_writes_handler", - IgnoreAllWritesHandler.INSTANCE) - .context("ignore_all_writes_handler"); - - segment.effectivelyBlocked(); - - ctxForDownstream - .writeAndFlush(response) + if (ctxForDownstream.pipeline().get(IGNORE_ALL_WRITES_HANDLER) == null) { + ctxForDownstream + .pipeline() + .addAfter( + ctxForDownstream.name(), IGNORE_ALL_WRITES_HANDLER, IgnoreAllWritesHandler.INSTANCE); + } + ChannelHandlerContext writeContext = + ctxForDownstream.pipeline().context(IGNORE_ALL_WRITES_HANDLER); + + writeContext + .writeAndFlush(pendingBlockResponse.toResponse()) .addListener( fut -> { if (!fut.isSuccess()) { log.warn("Write of blocking response failed", fut.cause()); } - ctx.channel().close(); + writeContext.channel().close(); }); } + private static class PendingBlockResponse { + private final TraceSegment segment; + private final int statusCode; + private final BlockingContentType bct; + private final Map extraHeaders; + private final String securityResponseId; + private final HttpVersion protocolVersion; + private final String acceptHeader; + + // Prevent the generation of BlockingResponseHandler$1 by making this constructor + // package-private to allow the BlockingResponseHandler to call this. This module emits Java 8 + // bytecode, so it cannot use Java 11 nested access. + PendingBlockResponse( + TraceSegment segment, + int statusCode, + BlockingContentType bct, + Map extraHeaders, + String securityResponseId, + HttpVersion protocolVersion, + String acceptHeader) { + this.segment = segment; + this.statusCode = statusCode; + this.bct = bct; + this.extraHeaders = extraHeaders; + this.securityResponseId = securityResponseId; + this.protocolVersion = protocolVersion; + this.acceptHeader = acceptHeader; + } + + FullHttpResponse toResponse() { + int httpCode = BlockingActionHelper.getHttpCode(statusCode); + HttpResponseStatus httpResponseStatus = HttpResponseStatus.valueOf(httpCode); + FullHttpResponse response = new DefaultFullHttpResponse(protocolVersion, httpResponseStatus); + + HttpHeaders headers = response.headers(); + headers.set("Connection", "close"); + + for (Map.Entry h : extraHeaders.entrySet()) { + headers.set(h.getKey(), h.getValue()); + } + + if (bct != BlockingContentType.NONE) { + BlockingActionHelper.TemplateType type = + BlockingActionHelper.determineTemplateType(bct, acceptHeader); + headers.set("Content-type", BlockingActionHelper.getContentType(type)); + + byte[] template = BlockingActionHelper.getTemplate(type, securityResponseId); + HttpUtil.setContentLength(response, template.length); + response.content().writeBytes(template); + } + segment.effectivelyBlocked(); + return response; + } + } + @ChannelHandler.Sharable public static class IgnoreAllWritesHandler extends ChannelOutboundHandlerAdapter { public static final IgnoreAllWritesHandler INSTANCE = new IgnoreAllWritesHandler(); diff --git a/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/server/HttpServerRequestTracingHandler.java b/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/server/HttpServerRequestTracingHandler.java index ce58fac2ded..71cf78d6a87 100644 --- a/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/server/HttpServerRequestTracingHandler.java +++ b/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/server/HttpServerRequestTracingHandler.java @@ -1,21 +1,20 @@ package datadog.trace.instrumentation.netty41.server; -import static datadog.trace.instrumentation.netty41.AttributeKeys.ANALYZED_RESPONSE_KEY; -import static datadog.trace.instrumentation.netty41.AttributeKeys.BLOCKED_RESPONSE_KEY; import static datadog.trace.instrumentation.netty41.AttributeKeys.CONTEXT_ATTRIBUTE_KEY; import static datadog.trace.instrumentation.netty41.AttributeKeys.PARENT_CONTEXT_ATTRIBUTE_KEY; -import static datadog.trace.instrumentation.netty41.AttributeKeys.REQUEST_HEADERS_ATTRIBUTE_KEY; import static datadog.trace.instrumentation.netty41.server.NettyHttpServerDecorator.DECORATE; import datadog.context.Context; import datadog.context.ContextScope; import datadog.trace.api.gateway.Flow; import datadog.trace.bootstrap.instrumentation.api.AgentSpan; +import datadog.trace.instrumentation.netty41.ServerRequestContext; import io.netty.channel.Channel; import io.netty.channel.ChannelHandler; import io.netty.channel.ChannelHandlerContext; import io.netty.channel.ChannelInboundHandlerAdapter; import io.netty.handler.codec.http.HttpHeaders; +import io.netty.handler.codec.http.HttpMethod; import io.netty.handler.codec.http.HttpRequest; @ChannelHandler.Sharable @@ -38,30 +37,34 @@ public void channelRead(final ChannelHandlerContext ctx, final Object msg) { } final HttpRequest request = (HttpRequest) msg; + if (!ServerRequestContext.canTrackRequest(channel)) { + channel.attr(PARENT_CONTEXT_ATTRIBUTE_KEY).remove(); + ctx.fireChannelRead(msg); + return; + } + final HttpHeaders headers = request.headers(); final Context storedParentContext = channel.attr(PARENT_CONTEXT_ATTRIBUTE_KEY).getAndRemove(); final Context parentContext = storedParentContext != null ? storedParentContext : DECORATE.extract(headers); final Context context = DECORATE.startSpan(headers, parentContext); + final ServerRequestContext serverContext = + ServerRequestContext.add( + channel, context, headers.get("accept"), HttpMethod.HEAD.equals(request.method())); try (final ContextScope ignored = context.attach()) { final AgentSpan span = AgentSpan.fromContext(context); DECORATE.afterStart(span); DECORATE.onRequest(span, channel, request, parentContext); - channel.attr(ANALYZED_RESPONSE_KEY).set(null); - channel.attr(BLOCKED_RESPONSE_KEY).set(null); - - channel.attr(CONTEXT_ATTRIBUTE_KEY).set(context); - channel.attr(REQUEST_HEADERS_ATTRIBUTE_KEY).set(request.headers()); - Flow.Action.RequestBlockingAction rba = span.getRequestBlockingAction(); if (rba != null) { ctx.pipeline() .addAfter( ctx.name(), "blocking_handler", - new BlockingResponseHandler(span.getRequestContext().getTraceSegment(), rba)); + new BlockingResponseHandler( + span.getRequestContext().getTraceSegment(), rba, serverContext)); } try { @@ -76,7 +79,7 @@ public void channelRead(final ChannelHandlerContext ctx, final Object msg) { DECORATE.onError(span, throwable); DECORATE.beforeFinish(ignored.context()); span.finish(); // Finish the span manually since finishSpanOnClose was false - ctx.channel().attr(CONTEXT_ATTRIBUTE_KEY).remove(); + ServerRequestContext.remove(ctx.channel(), serverContext); throw throwable; } } @@ -88,12 +91,7 @@ public void channelInactive(ChannelHandlerContext ctx) throws Exception { super.channelInactive(ctx); } finally { try { - final Context storedContext = ctx.channel().attr(CONTEXT_ATTRIBUTE_KEY).getAndRemove(); - final AgentSpan span = AgentSpan.fromContext(storedContext); - if (span != null && span.phasedFinish()) { - // at this point we can just publish this span to avoid loosing the rest of the trace - span.publish(); - } + ServerRequestContext.closeAll(ctx.channel()); } catch (final Throwable ignored) { } } diff --git a/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/server/HttpServerResponseTracingHandler.java b/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/server/HttpServerResponseTracingHandler.java index 04d7f8893fc..b45c3767066 100644 --- a/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/server/HttpServerResponseTracingHandler.java +++ b/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/server/HttpServerResponseTracingHandler.java @@ -8,13 +8,19 @@ import datadog.context.ContextScope; import datadog.trace.bootstrap.instrumentation.api.AgentSpan; import datadog.trace.bootstrap.instrumentation.websocket.HandlerContext; +import datadog.trace.instrumentation.netty41.ServerRequestContext; +import io.netty.buffer.Unpooled; import io.netty.channel.ChannelHandler; import io.netty.channel.ChannelHandlerContext; import io.netty.channel.ChannelOutboundHandlerAdapter; import io.netty.channel.ChannelPromise; +import io.netty.handler.codec.http.DefaultLastHttpContent; +import io.netty.handler.codec.http.FullHttpResponse; import io.netty.handler.codec.http.HttpHeaderNames; import io.netty.handler.codec.http.HttpResponse; import io.netty.handler.codec.http.HttpResponseStatus; +import io.netty.handler.codec.http.HttpUtil; +import io.netty.handler.codec.http.LastHttpContent; @ChannelHandler.Sharable public class HttpServerResponseTracingHandler extends ChannelOutboundHandlerAdapter { @@ -22,16 +28,35 @@ public class HttpServerResponseTracingHandler extends ChannelOutboundHandlerAdap @Override public void write(final ChannelHandlerContext ctx, final Object msg, final ChannelPromise prm) { - final Context storedContext = ctx.channel().attr(CONTEXT_ATTRIBUTE_KEY).get(); + final boolean isResponse = msg instanceof HttpResponse; + if (!isResponse && !(msg instanceof LastHttpContent)) { + ctx.write(msg, prm); + return; + } + + final ServerRequestContext serverContext = ServerRequestContext.nextResponse(ctx.channel()); + if (!isResponse && serverContext != null && !serverContext.isResponseStarted()) { + ctx.write(msg, prm); + return; + } + + final Context storedContext = + serverContext == null + // HTTP/2 multiplex stream channels only inherit the mirrored context attribute from + // Http2MultiplexHandlerStreamChannelInstrumentation.PropagateContextAdvice, without a + // per-stream request queue. + ? ctx.channel().attr(CONTEXT_ATTRIBUTE_KEY).get() + : serverContext.tracingContext(); final AgentSpan span = AgentSpan.fromContext(storedContext); - if (span == null || !(msg instanceof HttpResponse)) { + if (span == null) { ctx.write(msg, prm); return; } try (final ContextScope scope = storedContext.attach()) { - final HttpResponse response = (HttpResponse) msg; + final HttpResponse response = isResponse ? (HttpResponse) msg : null; + final boolean headerOnly = response != null && isHeaderOnly(response, serverContext); try { ctx.write(msg, prm); @@ -39,24 +64,75 @@ public void write(final ChannelHandlerContext ctx, final Object msg, final Chann DECORATE.onError(span, throwable); span.setHttpStatusCode(500); span.finish(); // Finish the span manually since finishSpanOnClose was false - ctx.channel().attr(CONTEXT_ATTRIBUTE_KEY).remove(); + removeServerContext(ctx, serverContext); throw throwable; } - final boolean isWebsocketUpgrade = - response.status() == HttpResponseStatus.SWITCHING_PROTOCOLS - && "websocket".equals(response.headers().get(HttpHeaderNames.UPGRADE)); - if (isWebsocketUpgrade) { - ctx.channel() - .attr(WEBSOCKET_SENDER_HANDLER_CONTEXT) - .set(new HandlerContext.Sender(span, ctx.channel().id().asShortText())); - } - if (response.status() != HttpResponseStatus.CONTINUE - && (response.status() != HttpResponseStatus.SWITCHING_PROTOCOLS || isWebsocketUpgrade)) { + final boolean responseComplete; + if (response == null) { + responseComplete = true; + } else { + final boolean isWebsocketUpgrade = + response.status() == HttpResponseStatus.SWITCHING_PROTOCOLS + && "websocket".equals(response.headers().get(HttpHeaderNames.UPGRADE)); + if (isWebsocketUpgrade) { + ctx.channel() + .attr(WEBSOCKET_SENDER_HANDLER_CONTEXT) + .set(new HandlerContext.Sender(span, ctx.channel().id().asShortText())); + } + if (isInformational(response) && !isWebsocketUpgrade) { + return; + } + if (serverContext != null) { + serverContext.markResponseStarted(); + } DECORATE.onResponse(span, response); + responseComplete = isWebsocketUpgrade || msg instanceof FullHttpResponse || headerOnly; + } + if (responseComplete) { DECORATE.beforeFinish(scope.context()); span.finish(); // Finish the span manually since finishSpanOnClose was false - ctx.channel().attr(CONTEXT_ATTRIBUTE_KEY).remove(); + removeServerContext(ctx, serverContext); + final ServerRequestContext nextResponse = ServerRequestContext.nextResponse(ctx.channel()); + if (headerOnly && !(msg instanceof FullHttpResponse) && nextResponse != null) { + ctx.write(new DefaultLastHttpContent(Unpooled.EMPTY_BUFFER), ctx.voidPromise()); + } + BlockingResponseHandler.maybeWriteDeferredBlockResponse(ctx, nextResponse); } } } + + private static void removeServerContext( + final ChannelHandlerContext ctx, final ServerRequestContext serverContext) { + if (serverContext == null) { + ctx.channel().attr(CONTEXT_ATTRIBUTE_KEY).remove(); + } else { + ServerRequestContext.remove(ctx.channel(), serverContext); + } + } + + private static boolean isHeaderOnly( + final HttpResponse response, final ServerRequestContext serverContext) { + final int statusCode = response.status().code(); + return (serverContext != null && serverContext.isHeadRequest()) + || statusCode == HttpResponseStatus.NO_CONTENT.code() + || statusCode == HttpResponseStatus.RESET_CONTENT.code() + || statusCode == HttpResponseStatus.NOT_MODIFIED.code() + || (hasZeroContentLength(response) && !HttpUtil.isTransferEncodingChunked(response)); + } + + private static boolean isInformational(final HttpResponse response) { + final int statusCode = response.status().code(); + return statusCode >= 100 && statusCode < 200; + } + + private static boolean hasZeroContentLength(final HttpResponse response) { + if (!HttpUtil.isContentLengthSet(response)) { + return false; + } + try { + return HttpUtil.getContentLength(response) == 0; + } catch (final NumberFormatException ignored) { + return false; + } + } } diff --git a/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/server/MaybeBlockResponseHandler.java b/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/server/MaybeBlockResponseHandler.java index 4f79bd6aea6..c9e18635dfe 100644 --- a/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/server/MaybeBlockResponseHandler.java +++ b/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/server/MaybeBlockResponseHandler.java @@ -1,9 +1,6 @@ package datadog.trace.instrumentation.netty41.server; -import static datadog.trace.instrumentation.netty41.AttributeKeys.ANALYZED_RESPONSE_KEY; -import static datadog.trace.instrumentation.netty41.AttributeKeys.BLOCKED_RESPONSE_KEY; import static datadog.trace.instrumentation.netty41.AttributeKeys.CONTEXT_ATTRIBUTE_KEY; -import static datadog.trace.instrumentation.netty41.AttributeKeys.REQUEST_HEADERS_ATTRIBUTE_KEY; import static datadog.trace.instrumentation.netty41.server.NettyHttpServerDecorator.DECORATE; import static io.netty.handler.codec.http.HttpHeaders.setContentLength; @@ -13,6 +10,7 @@ import datadog.trace.api.gateway.RequestContext; import datadog.trace.bootstrap.blocking.BlockingActionHelper; import datadog.trace.bootstrap.instrumentation.api.AgentSpan; +import datadog.trace.instrumentation.netty41.ServerRequestContext; import io.netty.channel.Channel; import io.netty.channel.ChannelHandler; import io.netty.channel.ChannelHandlerContext; @@ -25,6 +23,7 @@ import io.netty.handler.codec.http.HttpResponse; import io.netty.handler.codec.http.HttpResponseStatus; import io.netty.util.ReferenceCountUtil; +import java.nio.channels.ClosedChannelException; import java.util.Map; import org.slf4j.Logger; import org.slf4j.LoggerFactory; @@ -36,27 +35,23 @@ public class MaybeBlockResponseHandler extends ChannelOutboundHandlerAdapter { private MaybeBlockResponseHandler() {} - private static boolean isAnalyzedResponse(Channel ch) { - return ch.attr(ANALYZED_RESPONSE_KEY).get() != null; - } - - private static void markAnalyzedResponse(Channel ch) { - ch.attr(ANALYZED_RESPONSE_KEY).set(Boolean.TRUE); - } - - private static boolean isBlockedResponse(Channel ch) { - return ch.attr(BLOCKED_RESPONSE_KEY).get() != null; - } - - private static void markBlockedResponse(Channel ch) { - ch.attr(BLOCKED_RESPONSE_KEY).set(Boolean.TRUE); - } - @Override public void write(ChannelHandlerContext ctx, Object msg, ChannelPromise prm) throws Exception { Channel channel = ctx.channel(); - Context storedContext = channel.attr(CONTEXT_ATTRIBUTE_KEY).get(); + if (ServerRequestContext.isResponseBlocked(channel)) { + // block further writes while the blocking response close is still asynchronous + log.debug("Write suppressed; dropped outbound message"); + ReferenceCountUtil.release(msg); + prm.tryFailure(new ClosedChannelException()); + return; + } + + ServerRequestContext serverContext = ServerRequestContext.nextResponse(channel); + Context storedContext = + serverContext == null + ? channel.attr(CONTEXT_ATTRIBUTE_KEY).get() + : serverContext.tracingContext(); AgentSpan span = AgentSpan.fromContext(storedContext); RequestContext requestContext; if (span == null || (requestContext = span.getRequestContext()) == null) { @@ -64,14 +59,8 @@ public void write(ChannelHandlerContext ctx, Object msg, ChannelPromise prm) thr return; } - if (isAnalyzedResponse(channel)) { - if (isBlockedResponse(channel)) { - // block further writes - log.debug("Write suppressed, msg {} dropped", msg); - ReferenceCountUtil.release(msg); - } else { - super.write(ctx, msg, prm); - } + if (serverContext != null && serverContext.isResponseAnalyzed()) { + super.write(ctx, msg, prm); return; } @@ -80,22 +69,27 @@ public void write(ChannelHandlerContext ctx, Object msg, ChannelPromise prm) thr return; } HttpResponse origResponse = (HttpResponse) msg; - if (origResponse.status().code() == HttpResponseStatus.CONTINUE.code()) { + int statusCode = origResponse.status().code(); + if (statusCode >= 100 + && statusCode < 200 + && statusCode != HttpResponseStatus.SWITCHING_PROTOCOLS.code()) { super.write(ctx, msg, prm); return; } Flow flow = DECORATE.callIGCallbackResponseAndHeaders( - span, origResponse, origResponse.getStatus().code(), ResponseExtractAdapter.GETTER); - markAnalyzedResponse(channel); + span, origResponse, statusCode, ResponseExtractAdapter.GETTER); + if (serverContext != null) { + serverContext.markResponseAnalyzed(); + } Flow.Action action = flow.getAction(); if (!(action instanceof Flow.Action.RequestBlockingAction)) { super.write(ctx, msg, prm); return; } - markBlockedResponse(channel); + ServerRequestContext.markResponseBlocked(channel); Flow.Action.RequestBlockingAction rba = (Flow.Action.RequestBlockingAction) action; int httpCode = BlockingActionHelper.getHttpCode(rba.getStatusCode()); HttpResponseStatus httpResponseStatus = HttpResponseStatus.valueOf(httpCode); @@ -112,10 +106,9 @@ public void write(ChannelHandlerContext ctx, Object msg, ChannelPromise prm) thr BlockingContentType bct = rba.getBlockingContentType(); if (bct != BlockingContentType.NONE) { - HttpHeaders reqHeaders = ctx.attr(REQUEST_HEADERS_ATTRIBUTE_KEY).get(); - String acceptHeader = reqHeaders != null ? reqHeaders.get("accept") : null; BlockingActionHelper.TemplateType type = - BlockingActionHelper.determineTemplateType(bct, acceptHeader); + BlockingActionHelper.determineTemplateType( + bct, serverContext == null ? null : serverContext.acceptHeader()); headers.set("Content-type", BlockingActionHelper.getContentType(type)); byte[] template = BlockingActionHelper.getTemplate(type, rba.getSecurityResponseId()); setContentLength(response, template.length); diff --git a/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/server/NettyHttpServerDecorator.java b/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/server/NettyHttpServerDecorator.java index 101b35bc4ac..a66fda5f11c 100644 --- a/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/server/NettyHttpServerDecorator.java +++ b/dd-java-agent/instrumentation/netty/netty-4.1/src/main/java/datadog/trace/instrumentation/netty41/server/NettyHttpServerDecorator.java @@ -10,6 +10,7 @@ import datadog.trace.bootstrap.instrumentation.api.URIDefaultDataAdapter; import datadog.trace.bootstrap.instrumentation.api.UTF8BytesString; import datadog.trace.bootstrap.instrumentation.decorator.HttpServerDecorator; +import datadog.trace.instrumentation.netty41.ServerRequestContext; import io.netty.channel.Channel; import io.netty.channel.ChannelHandler; import io.netty.channel.ChannelHandlerContext; @@ -109,17 +110,24 @@ protected boolean isAppSecOnResponseSeparate() { @Override protected BlockResponseFunction createBlockResponseFunction( HttpRequest httpRequest, Channel channel) { - return new NettyBlockResponseFunction(channel.pipeline(), httpRequest); + return new NettyBlockResponseFunction( + channel.pipeline(), httpRequest, ServerRequestContext.currentRequest(channel)); } public static class NettyBlockResponseFunction implements BlockResponseFunction { - private final ChannelPipeline pipeline; public static final Logger log = LoggerFactory.getLogger(NettyBlockResponseFunction.class); + + private final ChannelPipeline pipeline; private final HttpRequest httpRequestMessage; + private final ServerRequestContext serverContext; - public NettyBlockResponseFunction(ChannelPipeline pipeline, HttpRequest httpRequestMessage) { + public NettyBlockResponseFunction( + ChannelPipeline pipeline, + HttpRequest httpRequestMessage, + ServerRequestContext serverContext) { this.pipeline = pipeline; this.httpRequestMessage = httpRequestMessage; + this.serverContext = serverContext; } @Override @@ -145,7 +153,12 @@ public boolean tryCommitBlockingResponse( pipeline.context(handlerBefore).name(), "blocking_handler", new BlockingResponseHandler( - segment, statusCode, templateType, extraHeaders, securityResponseId)) + segment, + statusCode, + templateType, + extraHeaders, + securityResponseId, + serverContext)) .addBefore( "blocking_handler", "before_blocking_handler", new ChannelInboundHandlerAdapter()); } catch (RuntimeException rte) { diff --git a/dd-java-agent/instrumentation/netty/netty-4.1/src/test/java/datadog/trace/instrumentation/netty41/ServerRequestContextTest.java b/dd-java-agent/instrumentation/netty/netty-4.1/src/test/java/datadog/trace/instrumentation/netty41/ServerRequestContextTest.java new file mode 100644 index 00000000000..684de404891 --- /dev/null +++ b/dd-java-agent/instrumentation/netty/netty-4.1/src/test/java/datadog/trace/instrumentation/netty41/ServerRequestContextTest.java @@ -0,0 +1,68 @@ +package datadog.trace.instrumentation.netty41; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import datadog.context.Context; +import io.netty.handler.codec.http.DefaultHttpHeaders; +import io.netty.util.DefaultAttributeMap; +import org.junit.jupiter.api.Test; + +class ServerRequestContextTest { + private static final int PIPELINING_LIMIT = 1000; + + @Test + void disablesTrackingWhenPipeliningLimitIsExceeded() { + DefaultAttributeMap attributes = new DefaultAttributeMap(); + + for (int i = 0; i < PIPELINING_LIMIT; i++) { + assertNotNull(ServerRequestContext.add(attributes, Context.root(), null)); + } + + assertFalse(ServerRequestContext.canTrackRequest(attributes)); + assertNull(ServerRequestContext.nextResponse(attributes)); + assertNull(attributes.attr(AttributeKeys.CONTEXT_ATTRIBUTE_KEY).get()); + assertNull(ServerRequestContext.add(attributes, Context.root(), null)); + } + + @Test + void capturesAcceptHeaderValue() { + DefaultAttributeMap attributes = new DefaultAttributeMap(); + DefaultHttpHeaders headers = new DefaultHttpHeaders(); + headers.set("accept", "text/html"); + + ServerRequestContext serverContext = + ServerRequestContext.add(attributes, Context.root(), headers.get("accept")); + headers.set("accept", "application/json"); + + assertEquals("text/html", serverContext.acceptHeader()); + } + + @Test + void capturesHeadRequest() { + DefaultAttributeMap attributes = new DefaultAttributeMap(); + + ServerRequestContext serverContext = + ServerRequestContext.add(attributes, Context.root(), null, true); + + assertTrue(serverContext.isHeadRequest()); + } + + @Test + void tracksBlockedResponseUntilChannelClose() { + DefaultAttributeMap attributes = new DefaultAttributeMap(); + ServerRequestContext serverContext = ServerRequestContext.add(attributes, Context.root(), null); + + ServerRequestContext.markResponseBlocked(attributes); + ServerRequestContext.remove(attributes, serverContext); + + assertTrue(ServerRequestContext.isResponseBlocked(attributes)); + + ServerRequestContext.closeAll(attributes); + + assertFalse(ServerRequestContext.isResponseBlocked(attributes)); + } +} diff --git a/dd-java-agent/instrumentation/netty/netty-4.1/src/test/java/datadog/trace/instrumentation/netty41/server/HttpServerResponseTracingHandlerTest.java b/dd-java-agent/instrumentation/netty/netty-4.1/src/test/java/datadog/trace/instrumentation/netty41/server/HttpServerResponseTracingHandlerTest.java new file mode 100644 index 00000000000..cb204d7ef51 --- /dev/null +++ b/dd-java-agent/instrumentation/netty/netty-4.1/src/test/java/datadog/trace/instrumentation/netty41/server/HttpServerResponseTracingHandlerTest.java @@ -0,0 +1,75 @@ +package datadog.trace.instrumentation.netty41.server; + +import static datadog.trace.agent.test.assertions.SpanMatcher.span; +import static datadog.trace.agent.test.assertions.TraceMatcher.trace; +import static datadog.trace.bootstrap.instrumentation.api.AgentTracer.startSpan; +import static datadog.trace.instrumentation.netty41.AttributeKeys.CONTEXT_ATTRIBUTE_KEY; +import static io.netty.handler.codec.http.HttpHeaderNames.CONTENT_LENGTH; +import static io.netty.handler.codec.http.HttpResponseStatus.OK; +import static io.netty.handler.codec.http.HttpVersion.HTTP_1_1; +import static org.junit.jupiter.api.Assertions.assertDoesNotThrow; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertSame; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import datadog.context.Context; +import datadog.trace.agent.test.AbstractInstrumentationTest; +import datadog.trace.bootstrap.instrumentation.api.AgentSpan; +import datadog.trace.instrumentation.netty41.ServerRequestContext; +import io.netty.channel.embedded.EmbeddedChannel; +import io.netty.handler.codec.http.DefaultFullHttpResponse; +import io.netty.handler.codec.http.DefaultLastHttpContent; +import io.netty.handler.codec.http.FullHttpResponse; +import io.netty.handler.codec.http.LastHttpContent; +import org.junit.jupiter.api.Test; + +class HttpServerResponseTracingHandlerTest extends AbstractInstrumentationTest { + + @Test + void finishesMirroredContextWhenRequestQueueIsAbsent() { + EmbeddedChannel channel = new EmbeddedChannel(HttpServerResponseTracingHandler.INSTANCE); + AgentSpan span = startSpan("netty", "mirrored-http2-server"); + channel.attr(CONTEXT_ATTRIBUTE_KEY).set(span); + + assertTrue(channel.writeOutbound(new DefaultFullHttpResponse(HTTP_1_1, OK))); + + FullHttpResponse response = channel.readOutbound(); + assertNotNull(response); + response.release(); + assertNull(channel.attr(CONTEXT_ATTRIBUTE_KEY).get()); + channel.finishAndReleaseAll(); + assertTraces(trace(span().root().operationName("mirrored-http2-server"))); + } + + @Test + void forwardsLastContentBeforeFinalResponseWhenItWasNotSyntheticTerminator() { + EmbeddedChannel channel = new EmbeddedChannel(HttpServerResponseTracingHandler.INSTANCE); + ServerRequestContext.add(channel, Context.root(), null); + LastHttpContent lastContent = new DefaultLastHttpContent(); + + assertTrue(channel.writeOutbound(lastContent)); + + LastHttpContent forwarded = channel.readOutbound(); + assertSame(lastContent, forwarded); + forwarded.release(); + channel.finishAndReleaseAll(); + } + + @Test + void doesNotThrowOnMalformedContentLength() { + EmbeddedChannel channel = new EmbeddedChannel(HttpServerResponseTracingHandler.INSTANCE); + AgentSpan span = startSpan("netty", "malformed-content-length-server"); + ServerRequestContext.add(channel, span, null); + FullHttpResponse response = new DefaultFullHttpResponse(HTTP_1_1, OK); + response.headers().set(CONTENT_LENGTH, "malformed"); + + assertDoesNotThrow(() -> assertTrue(channel.writeOutbound(response))); + + FullHttpResponse forwarded = channel.readOutbound(); + assertSame(response, forwarded); + forwarded.release(); + channel.finishAndReleaseAll(); + assertTraces(trace(span().root().operationName("malformed-content-length-server"))); + } +} diff --git a/dd-java-agent/instrumentation/netty/netty-4.1/src/test/java/datadog/trace/instrumentation/netty41/server/MaybeBlockResponseHandlerTest.java b/dd-java-agent/instrumentation/netty/netty-4.1/src/test/java/datadog/trace/instrumentation/netty41/server/MaybeBlockResponseHandlerTest.java new file mode 100644 index 00000000000..2f98851a807 --- /dev/null +++ b/dd-java-agent/instrumentation/netty/netty-4.1/src/test/java/datadog/trace/instrumentation/netty41/server/MaybeBlockResponseHandlerTest.java @@ -0,0 +1,172 @@ +package datadog.trace.instrumentation.netty41.server; + +import static datadog.trace.api.gateway.Events.EVENTS; +import static datadog.trace.instrumentation.netty41.AttributeKeys.CONTEXT_ATTRIBUTE_KEY; +import static datadog.trace.instrumentation.netty41.server.NettyHttpServerDecorator.DECORATE; +import static io.netty.handler.codec.http.HttpHeaderNames.UPGRADE; +import static io.netty.handler.codec.http.HttpResponseStatus.FORBIDDEN; +import static io.netty.handler.codec.http.HttpResponseStatus.OK; +import static io.netty.handler.codec.http.HttpResponseStatus.SWITCHING_PROTOCOLS; +import static io.netty.handler.codec.http.HttpVersion.HTTP_1_1; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import datadog.appsec.api.blocking.BlockingContentType; +import datadog.context.Context; +import datadog.trace.agent.test.AbstractInstrumentationTest; +import datadog.trace.api.function.TriConsumer; +import datadog.trace.api.gateway.Flow; +import datadog.trace.api.gateway.RequestContext; +import datadog.trace.api.gateway.RequestContextSlot; +import datadog.trace.api.gateway.SubscriptionService; +import datadog.trace.bootstrap.ActiveSubsystems; +import datadog.trace.bootstrap.instrumentation.api.AgentSpan; +import datadog.trace.bootstrap.instrumentation.api.AgentTracer; +import datadog.trace.instrumentation.netty41.ServerRequestContext; +import io.netty.buffer.ByteBuf; +import io.netty.buffer.Unpooled; +import io.netty.channel.ChannelPromise; +import io.netty.channel.embedded.EmbeddedChannel; +import io.netty.handler.codec.http.DefaultFullHttpResponse; +import io.netty.handler.codec.http.DefaultHttpHeaders; +import io.netty.handler.codec.http.FullHttpResponse; +import io.netty.handler.codec.http.HttpResponseStatus; +import io.netty.util.ReferenceCountUtil; +import java.nio.channels.ClosedChannelException; +import java.util.function.Function; +import java.util.function.Supplier; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.Test; + +class MaybeBlockResponseHandlerTest extends AbstractInstrumentationTest { + + private static final HttpResponseStatus EARLY_HINTS = new HttpResponseStatus(103, "Early Hints"); + + private Object appSecSubscriptions; + private boolean originalAppSecActive; + + @AfterEach + void resetAppSec() { + if (appSecSubscriptions != null) { + ((SubscriptionService) appSecSubscriptions).reset(); + appSecSubscriptions = null; + ActiveSubsystems.APPSEC_ACTIVE = originalAppSecActive; + } + } + + @Test + void blocksFinalResponseUsingMirroredContextAfterInformationalResponse() { + enableAppSecResponseBlocking(); + Context context = DECORATE.startSpan(new DefaultHttpHeaders(), Context.root()); + AgentSpan span = AgentSpan.fromContext(context); + EmbeddedChannel channel = new EmbeddedChannel(MaybeBlockResponseHandler.INSTANCE); + channel.attr(CONTEXT_ATTRIBUTE_KEY).set(context); + FullHttpResponse informationalResponse = null; + FullHttpResponse response = null; + + try { + channel.writeOutbound(new DefaultFullHttpResponse(HTTP_1_1, EARLY_HINTS)); + + informationalResponse = channel.readOutbound(); + assertNotNull(informationalResponse); + assertEquals(EARLY_HINTS, informationalResponse.status()); + informationalResponse.release(); + informationalResponse = null; + + channel.writeOutbound(new DefaultFullHttpResponse(HTTP_1_1, OK)); + + response = channel.readOutbound(); + assertNotNull(response); + assertEquals(FORBIDDEN, response.status()); + } finally { + ReferenceCountUtil.release(informationalResponse); + ReferenceCountUtil.release(response); + channel.finishAndReleaseAll(); + span.finish(); + } + } + + @Test + void blocksNonWebSocketSwitchingProtocolsResponse() { + enableAppSecResponseBlocking(); + Context context = DECORATE.startSpan(new DefaultHttpHeaders(), Context.root()); + AgentSpan span = AgentSpan.fromContext(context); + EmbeddedChannel channel = new EmbeddedChannel(MaybeBlockResponseHandler.INSTANCE); + channel.attr(CONTEXT_ATTRIBUTE_KEY).set(context); + FullHttpResponse response = null; + + try { + FullHttpResponse switchingProtocols = + new DefaultFullHttpResponse(HTTP_1_1, SWITCHING_PROTOCOLS); + switchingProtocols.headers().set(UPGRADE, "h2c"); + + channel.writeOutbound(switchingProtocols); + + response = channel.readOutbound(); + assertNotNull(response); + assertEquals(FORBIDDEN, response.status()); + } finally { + ReferenceCountUtil.release(response); + channel.finishAndReleaseAll(); + span.finish(); + } + } + + @Test + void dropsWritesAfterBlockedContextHasBeenRemoved() { + EmbeddedChannel channel = new EmbeddedChannel(MaybeBlockResponseHandler.INSTANCE); + ServerRequestContext serverContext = ServerRequestContext.add(channel, Context.root(), null); + ServerRequestContext.markResponseBlocked(channel); + ServerRequestContext.remove(channel, serverContext); + ByteBuf lateResponseChunk = Unpooled.buffer().writeByte(1); + ChannelPromise promise = channel.newPromise(); + + channel.pipeline().write(lateResponseChunk, promise); + + assertEquals(0, lateResponseChunk.refCnt()); + assertTrue(promise.isDone()); + assertFalse(promise.isSuccess()); + assertTrue(promise.cause() instanceof ClosedChannelException); + assertNull(channel.readOutbound()); + channel.finishAndReleaseAll(); + } + + private void enableAppSecResponseBlocking() { + SubscriptionService subscriptions = + (SubscriptionService) AgentTracer.get().getSubscriptionService(RequestContextSlot.APPSEC); + appSecSubscriptions = subscriptions; + originalAppSecActive = ActiveSubsystems.APPSEC_ACTIVE; + ActiveSubsystems.APPSEC_ACTIVE = true; + + subscriptions.registerCallback( + EVENTS.requestStarted(), + new Supplier>() { + @Override + public Flow get() { + return new Flow.ResultFlow<>(new Object()); + } + }); + subscriptions.registerCallback( + EVENTS.responseHeader(), + new TriConsumer() { + @Override + public void accept(RequestContext requestContext, String name, String value) {} + }); + subscriptions.registerCallback( + EVENTS.responseHeaderDone(), + new Function>() { + @Override + public Flow apply(RequestContext requestContext) { + return new Flow.ResultFlow(null) { + @Override + public Action getAction() { + return new Action.RequestBlockingAction(403, BlockingContentType.AUTO); + } + }; + } + }); + } +} diff --git a/dd-java-agent/instrumentation/netty/netty-4.1/src/test/java/datadog/trace/instrumentation/netty41/server/NettyHttp11PipeliningTest.java b/dd-java-agent/instrumentation/netty/netty-4.1/src/test/java/datadog/trace/instrumentation/netty41/server/NettyHttp11PipeliningTest.java new file mode 100644 index 00000000000..05d678fc99f --- /dev/null +++ b/dd-java-agent/instrumentation/netty/netty-4.1/src/test/java/datadog/trace/instrumentation/netty41/server/NettyHttp11PipeliningTest.java @@ -0,0 +1,714 @@ +package datadog.trace.instrumentation.netty41.server; + +import static datadog.trace.agent.test.assertions.SpanMatcher.span; +import static datadog.trace.agent.test.assertions.TraceAssertions.SORT_BY_START_TIME; +import static datadog.trace.agent.test.assertions.TraceMatcher.trace; +import static datadog.trace.api.gateway.Events.EVENTS; +import static io.netty.handler.codec.http.HttpHeaderNames.CONNECTION; +import static io.netty.handler.codec.http.HttpHeaderNames.CONTENT_LENGTH; +import static io.netty.handler.codec.http.HttpHeaderNames.TRANSFER_ENCODING; +import static io.netty.handler.codec.http.HttpHeaderValues.CHUNKED; +import static io.netty.handler.codec.http.HttpHeaderValues.KEEP_ALIVE; +import static io.netty.handler.codec.http.HttpResponseStatus.CONTINUE; +import static io.netty.handler.codec.http.HttpResponseStatus.NO_CONTENT; +import static io.netty.handler.codec.http.HttpResponseStatus.OK; +import static io.netty.handler.codec.http.HttpVersion.HTTP_1_1; +import static java.nio.charset.StandardCharsets.US_ASCII; +import static java.nio.charset.StandardCharsets.UTF_8; +import static java.util.Collections.emptyMap; +import static java.util.concurrent.TimeUnit.SECONDS; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import datadog.appsec.api.blocking.BlockingContentType; +import datadog.trace.agent.test.AbstractInstrumentationTest; +import datadog.trace.agent.test.assertions.TraceMatcher; +import datadog.trace.api.function.TriFunction; +import datadog.trace.api.gateway.BlockResponseFunction; +import datadog.trace.api.gateway.Flow; +import datadog.trace.api.gateway.RequestContext; +import datadog.trace.api.gateway.RequestContextSlot; +import datadog.trace.api.gateway.SubscriptionService; +import datadog.trace.bootstrap.ActiveSubsystems; +import datadog.trace.bootstrap.instrumentation.api.AgentSpan; +import datadog.trace.bootstrap.instrumentation.api.AgentTracer; +import datadog.trace.bootstrap.instrumentation.api.URIDataAdapter; +import io.netty.bootstrap.ServerBootstrap; +import io.netty.buffer.Unpooled; +import io.netty.channel.Channel; +import io.netty.channel.ChannelHandler; +import io.netty.channel.ChannelHandlerContext; +import io.netty.channel.ChannelInitializer; +import io.netty.channel.EventLoopGroup; +import io.netty.channel.SimpleChannelInboundHandler; +import io.netty.channel.nio.NioEventLoopGroup; +import io.netty.channel.socket.nio.NioServerSocketChannel; +import io.netty.handler.codec.http.DefaultFullHttpResponse; +import io.netty.handler.codec.http.DefaultHttpContent; +import io.netty.handler.codec.http.DefaultHttpResponse; +import io.netty.handler.codec.http.DefaultLastHttpContent; +import io.netty.handler.codec.http.FullHttpRequest; +import io.netty.handler.codec.http.HttpObjectAggregator; +import io.netty.handler.codec.http.HttpResponseStatus; +import io.netty.handler.codec.http.HttpServerCodec; +import io.netty.util.ReferenceCountUtil; +import java.io.ByteArrayOutputStream; +import java.io.EOFException; +import java.io.IOException; +import java.io.InputStream; +import java.net.InetSocketAddress; +import java.net.Socket; +import java.util.ArrayList; +import java.util.List; +import java.util.concurrent.CountDownLatch; +import java.util.function.Supplier; +import java.util.regex.Pattern; +import org.junit.jupiter.api.AfterAll; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeAll; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.TestInstance; + +@TestInstance(TestInstance.Lifecycle.PER_CLASS) +public class NettyHttp11PipeliningTest extends AbstractInstrumentationTest { + + private static final String FIRST_PATH = "/pipelined/first"; + private static final String SECOND_PATH = "/pipelined/second"; + private static final String THIRD_PATH = "/pipelined/third"; + private static final HttpResponseStatus EARLY_HINTS = new HttpResponseStatus(103, "Early Hints"); + + private EventLoopGroup eventLoopGroup; + private PipeliningHandler handler; + private int port; + private Object appSecSubscriptions; + private boolean originalAppSecActive; + + @BeforeAll + void startServer() throws Exception { + eventLoopGroup = new NioEventLoopGroup(); + handler = new PipeliningHandler(); + ServerBootstrap bootstrap = + new ServerBootstrap() + .group(eventLoopGroup) + .channel(NioServerSocketChannel.class) + .childHandler( + new ChannelInitializer() { + @Override + protected void initChannel(Channel ch) { + ch.pipeline().addLast(new HttpServerCodec()); + ch.pipeline().addLast(new HttpObjectAggregator(65536)); + ch.pipeline().addLast(handler); + } + }); + Channel channel = bootstrap.bind(0).sync().channel(); + port = ((InetSocketAddress) channel.localAddress()).getPort(); + } + + @AfterAll + void stopServer() { + if (eventLoopGroup != null) { + eventLoopGroup.shutdownGracefully(); + } + } + + @AfterEach + void resetAppSec() { + if (appSecSubscriptions != null) { + ((SubscriptionService) appSecSubscriptions).reset(); + appSecSubscriptions = null; + ActiveSubsystems.APPSEC_ACTIVE = originalAppSecActive; + } + handler.clearBlockResponseFunctionBlocking(); + } + + @Test + void createsServerSpanForEachPipelinedRequest() throws Exception { + handler.expectRequests(3); + try (Socket socket = new Socket("localhost", port)) { + socket.setSoTimeout(5000); + socket.getOutputStream().write(pipelinedRequests().getBytes(US_ASCII)); + socket.getOutputStream().flush(); + + assertTrue( + handler.awaitAllRequestsReceived(), + "server did not receive all pipelined requests before responding"); + + handler.writeResponses(); + + assertEquals("response " + FIRST_PATH, readHttpResponseBody(socket.getInputStream())); + assertEquals("response " + SECOND_PATH, readHttpResponseBody(socket.getInputStream())); + assertEquals("response " + THIRD_PATH, readHttpResponseBody(socket.getInputStream())); + } + + assertTraces( + SORT_BY_START_TIME, + serverTrace(FIRST_PATH), + serverTrace(SECOND_PATH), + serverTrace(THIRD_PATH)); + } + + @Test + void requestBlockOnLaterPipelinedRequestDoesNotOvertakeEarlierResponse() throws Exception { + handler.expectRequests(1); + + try (Socket socket = new Socket("localhost", port)) { + socket.setSoTimeout(5000); + socket.getOutputStream().write(request(FIRST_PATH).getBytes(US_ASCII)); + socket.getOutputStream().flush(); + + assertTrue(handler.awaitAllRequestsReceived(), "server did not receive first request"); + + CountDownLatch blockedRequestSeen = enableAppSecRequestBlockingFor(SECOND_PATH); + socket.getOutputStream().write(request(SECOND_PATH).getBytes(US_ASCII)); + socket.getOutputStream().flush(); + + assertTrue( + blockedRequestSeen.await(5, SECONDS), "server did not block second pipelined request"); + + handler.writeResponses(); + + assertEquals("response " + FIRST_PATH, readHttpResponseBody(socket.getInputStream())); + assertTrue( + readHttpResponseHeaders(socket.getInputStream()).startsWith("HTTP/1.1 403 "), + "second response should be the deferred blocking response"); + } + } + + @Test + void requestBlockOnLaterPipelinedRequestWaitsForEarlierChunkedResponseCompletion() + throws Exception { + handler.expectRequests(1); + + try (Socket socket = new Socket("localhost", port)) { + socket.setSoTimeout(5000); + socket.getOutputStream().write(request(FIRST_PATH).getBytes(US_ASCII)); + socket.getOutputStream().flush(); + + assertTrue(handler.awaitAllRequestsReceived(), "server did not receive first request"); + + CountDownLatch blockedRequestSeen = enableAppSecRequestBlockingFor(SECOND_PATH); + socket.getOutputStream().write(request(SECOND_PATH).getBytes(US_ASCII)); + socket.getOutputStream().flush(); + + assertTrue( + blockedRequestSeen.await(5, SECONDS), "server did not block second pipelined request"); + + handler.writeChunkedResponse(); + + assertEquals("response " + FIRST_PATH, readHttpChunkedResponseBody(socket.getInputStream())); + assertTrue( + readHttpResponseHeaders(socket.getInputStream()).startsWith("HTTP/1.1 403 "), + "second response should be the deferred blocking response"); + } + } + + @Test + void requestBlockOnLaterPipelinedRequestFollowsEarlierHeaderOnlyResponse() throws Exception { + handler.expectRequests(1); + + try (Socket socket = new Socket("localhost", port)) { + socket.setSoTimeout(5000); + socket.getOutputStream().write(request(FIRST_PATH).getBytes(US_ASCII)); + socket.getOutputStream().flush(); + + assertTrue(handler.awaitAllRequestsReceived(), "server did not receive first request"); + + CountDownLatch blockedRequestSeen = enableAppSecRequestBlockingFor(SECOND_PATH); + socket.getOutputStream().write(request(SECOND_PATH).getBytes(US_ASCII)); + socket.getOutputStream().flush(); + + assertTrue( + blockedRequestSeen.await(5, SECONDS), "server did not block second pipelined request"); + + handler.writeHeaderOnlyResponse(); + + assertTrue( + readHttpResponseHeaders(socket.getInputStream()).startsWith("HTTP/1.1 204 "), + "first response should be the header-only response"); + assertTrue( + readHttpResponseHeaders(socket.getInputStream()).startsWith("HTTP/1.1 403 "), + "second response should be the deferred blocking response"); + } + } + + @Test + void requestBlockOnLaterPipelinedRequestFollowsEarlierHeadResponse() throws Exception { + handler.expectRequests(1); + + try (Socket socket = new Socket("localhost", port)) { + socket.setSoTimeout(5000); + socket.getOutputStream().write(headRequest(FIRST_PATH).getBytes(US_ASCII)); + socket.getOutputStream().flush(); + + assertTrue(handler.awaitAllRequestsReceived(), "server did not receive first request"); + + CountDownLatch blockedRequestSeen = enableAppSecRequestBlockingFor(SECOND_PATH); + socket.getOutputStream().write(request(SECOND_PATH).getBytes(US_ASCII)); + socket.getOutputStream().flush(); + + assertTrue( + blockedRequestSeen.await(5, SECONDS), "server did not block second pipelined request"); + + handler.writeHeadResponse(); + + assertTrue( + readHttpResponseHeaders(socket.getInputStream()).startsWith("HTTP/1.1 200 "), + "first response should be the HEAD response"); + assertTrue( + readHttpResponseHeaders(socket.getInputStream()).startsWith("HTTP/1.1 403 "), + "second response should be the deferred blocking response"); + } + } + + @Test + void lastContentAfterInterimResponseDoesNotCompleteServerSpan() throws Exception { + handler.expectRequests(1); + + try (Socket socket = new Socket("localhost", port)) { + socket.setSoTimeout(5000); + socket.getOutputStream().write(request(FIRST_PATH).getBytes(US_ASCII)); + socket.getOutputStream().flush(); + + assertTrue(handler.awaitAllRequestsReceived(), "server did not receive first request"); + + handler.writeInterimResponseWithTerminatorThenResponse(); + + assertTrue( + readHttpResponseHeaders(socket.getInputStream()).startsWith("HTTP/1.1 100 "), + "first response should be the interim response"); + assertEquals("response " + FIRST_PATH, readHttpResponseBody(socket.getInputStream())); + } + } + + @Test + void requestBlockOnLaterPipelinedRequestWaitsForEarlierEarlyHintsResponseCompletion() + throws Exception { + handler.expectRequests(1); + + try (Socket socket = new Socket("localhost", port)) { + socket.setSoTimeout(5000); + socket.getOutputStream().write(request(FIRST_PATH).getBytes(US_ASCII)); + socket.getOutputStream().flush(); + + assertTrue(handler.awaitAllRequestsReceived(), "server did not receive first request"); + + CountDownLatch blockedRequestSeen = enableAppSecRequestBlockingFor(SECOND_PATH); + socket.getOutputStream().write(request(SECOND_PATH).getBytes(US_ASCII)); + socket.getOutputStream().flush(); + + assertTrue( + blockedRequestSeen.await(5, SECONDS), "server did not block second pipelined request"); + + handler.writeEarlyHintsWithTerminatorThenResponse(); + + assertTrue( + readHttpResponseHeaders(socket.getInputStream()).startsWith("HTTP/1.1 103 "), + "first response should be the early hints response"); + assertEquals("response " + FIRST_PATH, readHttpResponseBody(socket.getInputStream())); + assertTrue( + readHttpResponseHeaders(socket.getInputStream()).startsWith("HTTP/1.1 403 "), + "second response should be the deferred blocking response"); + } + } + + @Test + void blockResponseFunctionOnLaterPipelinedRequestDoesNotOvertakeEarlierResponse() + throws Exception { + enableAppSec(); + handler.expectRequests(1); + + try (Socket socket = new Socket("localhost", port)) { + socket.setSoTimeout(5000); + socket.getOutputStream().write(request(FIRST_PATH).getBytes(US_ASCII)); + socket.getOutputStream().flush(); + + assertTrue(handler.awaitAllRequestsReceived(), "server did not receive first request"); + + CountDownLatch blockedRequestSeen = handler.blockWithResponseFunctionFor(SECOND_PATH); + socket.getOutputStream().write(request(SECOND_PATH).getBytes(US_ASCII)); + socket.getOutputStream().flush(); + + assertTrue( + blockedRequestSeen.await(5, SECONDS), "server did not block second pipelined request"); + + handler.writeResponses(); + + assertEquals("response " + FIRST_PATH, readHttpResponseBody(socket.getInputStream())); + assertTrue( + readHttpResponseHeaders(socket.getInputStream()).startsWith("HTTP/1.1 403 "), + "second response should be the deferred blocking response"); + } + } + + private static String pipelinedRequests() { + return "GET " + + FIRST_PATH + + " HTTP/1.1\r\nHost: localhost\r\n\r\n" + + "GET " + + SECOND_PATH + + " HTTP/1.1\r\nHost: localhost\r\n\r\n" + + "GET " + + THIRD_PATH + + " HTTP/1.1\r\nHost: localhost\r\n\r\n"; + } + + private static String request(String path) { + return "GET " + path + " HTTP/1.1\r\nHost: localhost\r\n\r\n"; + } + + private static String headRequest(String path) { + return "HEAD " + path + " HTTP/1.1\r\nHost: localhost\r\n\r\n"; + } + + private CountDownLatch enableAppSecRequestBlockingFor(String blockedPath) { + SubscriptionService subscriptions = (SubscriptionService) enableAppSec(); + CountDownLatch blockedRequestSeen = new CountDownLatch(1); + + subscriptions.registerCallback( + EVENTS.requestMethodUriRaw(), + new TriFunction>() { + @Override + public Flow apply( + RequestContext requestContext, String method, URIDataAdapter uri) { + if (!blockedPath.equals(uri.path())) { + return Flow.ResultFlow.empty(); + } + blockedRequestSeen.countDown(); + return new Flow.ResultFlow(null) { + @Override + public Action getAction() { + return new Action.RequestBlockingAction(403, BlockingContentType.NONE); + } + }; + } + }); + return blockedRequestSeen; + } + + private Object enableAppSec() { + SubscriptionService subscriptions = + (SubscriptionService) AgentTracer.get().getSubscriptionService(RequestContextSlot.APPSEC); + appSecSubscriptions = subscriptions; + originalAppSecActive = ActiveSubsystems.APPSEC_ACTIVE; + ActiveSubsystems.APPSEC_ACTIVE = true; + + subscriptions.registerCallback( + EVENTS.requestStarted(), + new Supplier>() { + @Override + public Flow get() { + return new Flow.ResultFlow<>(new Object()); + } + }); + return subscriptions; + } + + private static TraceMatcher serverTrace(String path) { + return trace( + span() + .root() + .operationName(Pattern.compile("netty\\.request")) + .resourceName(Pattern.compile("GET " + Pattern.quote(path))) + .type("web")); + } + + private static String readHttpResponseBody(InputStream in) throws IOException { + String headers = readHttpResponseHeaders(in); + assertTrue(headers.startsWith("HTTP/1.1 200 "), "unexpected response: " + headers); + int contentLength = contentLength(headers); + byte[] body = new byte[contentLength]; + int read = 0; + while (read < contentLength) { + int count = in.read(body, read, contentLength - read); + if (count == -1) { + throw new EOFException("response ended before body was complete"); + } + read += count; + } + return new String(body, UTF_8); + } + + private static String readHttpChunkedResponseBody(InputStream in) throws IOException { + String headers = readHttpResponseHeaders(in); + assertTrue(headers.startsWith("HTTP/1.1 200 "), "unexpected response: " + headers); + + ByteArrayOutputStream body = new ByteArrayOutputStream(); + while (true) { + String chunkSizeLine = readHttpLine(in); + int chunkSize = Integer.parseInt(chunkSizeLine, 16); + if (chunkSize == 0) { + String trailer; + do { + trailer = readHttpLine(in); + } while (!trailer.isEmpty()); + return body.toString(UTF_8.name()); + } + byte[] chunk = new byte[chunkSize]; + int read = 0; + while (read < chunkSize) { + int count = in.read(chunk, read, chunkSize - read); + if (count == -1) { + throw new EOFException("response ended before chunk was complete"); + } + read += count; + } + body.write(chunk); + String chunkTerminator = readHttpLine(in); + assertEquals("", chunkTerminator, "chunk was not followed by CRLF"); + } + } + + private static String readHttpResponseHeaders(InputStream in) throws IOException { + ByteArrayOutputStream out = new ByteArrayOutputStream(); + int state = 0; + while (state < 4) { + int b = in.read(); + if (b == -1) { + throw new EOFException("response ended before headers were complete"); + } + out.write(b); + if ((state == 0 || state == 2) && b == '\r') { + state++; + } else if ((state == 1 || state == 3) && b == '\n') { + state++; + } else { + state = b == '\r' ? 1 : 0; + } + } + return out.toString(US_ASCII.name()); + } + + private static String readHttpLine(InputStream in) throws IOException { + ByteArrayOutputStream out = new ByteArrayOutputStream(); + int previous = -1; + while (true) { + int current = in.read(); + if (current == -1) { + throw new EOFException("response ended before line was complete"); + } + if (previous == '\r' && current == '\n') { + byte[] line = out.toByteArray(); + return new String(line, 0, line.length - 1, US_ASCII); + } + out.write(current); + previous = current; + } + } + + private static int contentLength(String headers) { + for (String line : headers.split("\r\n")) { + int separator = line.indexOf(':'); + if (separator > 0 && "content-length".equalsIgnoreCase(line.substring(0, separator))) { + return Integer.parseInt(line.substring(separator + 1).trim()); + } + } + throw new AssertionError("missing content-length header: " + headers); + } + + @ChannelHandler.Sharable + private static final class PipeliningHandler + extends SimpleChannelInboundHandler { + private volatile CountDownLatch receivedRequests; + private final List paths = new ArrayList<>(); + private volatile ChannelHandlerContext context; + private volatile String blockResponseFunctionPath; + private volatile CountDownLatch blockResponseFunctionRequestSeen; + + private PipeliningHandler() { + super(false); + expectRequests(0); + } + + private void expectRequests(int expectedRequests) { + receivedRequests = new CountDownLatch(expectedRequests); + context = null; + synchronized (paths) { + paths.clear(); + } + } + + private CountDownLatch blockWithResponseFunctionFor(String path) { + blockResponseFunctionPath = path; + blockResponseFunctionRequestSeen = new CountDownLatch(1); + return blockResponseFunctionRequestSeen; + } + + private void clearBlockResponseFunctionBlocking() { + blockResponseFunctionPath = null; + blockResponseFunctionRequestSeen = null; + } + + @Override + protected void channelRead0(ChannelHandlerContext ctx, FullHttpRequest request) { + context = ctx; + boolean blockingResponseCommitted = false; + synchronized (paths) { + paths.add(request.uri()); + } + if (request.uri().equals(blockResponseFunctionPath)) { + AgentSpan span = AgentTracer.activeSpan(); + RequestContext requestContext = span == null ? null : span.getRequestContext(); + BlockResponseFunction blockResponseFunction = + requestContext == null ? null : requestContext.getBlockResponseFunction(); + if (blockResponseFunction != null) { + blockingResponseCommitted = + blockResponseFunction.tryCommitBlockingResponse( + requestContext.getTraceSegment(), + 403, + BlockingContentType.NONE, + emptyMap(), + null); + } + if (blockingResponseCommitted && blockResponseFunctionRequestSeen != null) { + blockResponseFunctionRequestSeen.countDown(); + } + } + receivedRequests.countDown(); + if (!blockingResponseCommitted) { + ReferenceCountUtil.release(request); + } + } + + private boolean awaitAllRequestsReceived() throws InterruptedException { + return receivedRequests.await(5, SECONDS); + } + + private void writeResponses() { + ChannelHandlerContext responseContext = context; + if (responseContext == null) { + throw new IllegalStateException("no request context captured"); + } + List responsePaths; + synchronized (paths) { + responsePaths = new ArrayList<>(paths); + } + responseContext + .executor() + .execute( + () -> { + for (String path : responsePaths) { + byte[] body = ("response " + path).getBytes(UTF_8); + DefaultFullHttpResponse response = + new DefaultFullHttpResponse(HTTP_1_1, OK, Unpooled.wrappedBuffer(body)); + response.headers().set(CONTENT_LENGTH, body.length); + responseContext.write(response); + } + responseContext.flush(); + }); + } + + private void writeChunkedResponse() { + ChannelHandlerContext responseContext = context; + if (responseContext == null) { + throw new IllegalStateException("no request context captured"); + } + String path; + synchronized (paths) { + path = paths.get(0); + } + responseContext + .executor() + .execute( + () -> { + byte[] body = ("response " + path).getBytes(UTF_8); + DefaultHttpResponse response = new DefaultHttpResponse(HTTP_1_1, OK); + response.headers().set(TRANSFER_ENCODING, CHUNKED); + responseContext.write(response); + responseContext.write(new DefaultHttpContent(Unpooled.wrappedBuffer(body))); + responseContext.write(new DefaultLastHttpContent(Unpooled.EMPTY_BUFFER)); + responseContext.flush(); + }); + } + + private void writeHeaderOnlyResponse() { + ChannelHandlerContext responseContext = context; + if (responseContext == null) { + throw new IllegalStateException("no request context captured"); + } + responseContext + .executor() + .execute( + () -> { + DefaultHttpResponse response = new DefaultHttpResponse(HTTP_1_1, NO_CONTENT); + response.headers().set(CONNECTION, KEEP_ALIVE); + responseContext.write(response); + responseContext.flush(); + }); + } + + private void writeHeadResponse() { + ChannelHandlerContext responseContext = context; + if (responseContext == null) { + throw new IllegalStateException("no request context captured"); + } + String path; + synchronized (paths) { + path = paths.get(0); + } + responseContext + .executor() + .execute( + () -> { + byte[] body = ("response " + path).getBytes(UTF_8); + DefaultHttpResponse response = new DefaultHttpResponse(HTTP_1_1, OK); + response.headers().set(CONTENT_LENGTH, body.length); + responseContext.write(response); + responseContext.flush(); + }); + } + + private void writeInterimResponseWithTerminatorThenResponse() { + ChannelHandlerContext responseContext = context; + if (responseContext == null) { + throw new IllegalStateException("no request context captured"); + } + String path; + synchronized (paths) { + path = paths.get(0); + } + responseContext + .executor() + .execute( + () -> { + responseContext.write(new DefaultHttpResponse(HTTP_1_1, CONTINUE)); + responseContext.write(new DefaultLastHttpContent(Unpooled.EMPTY_BUFFER)); + + byte[] body = ("response " + path).getBytes(UTF_8); + DefaultFullHttpResponse response = + new DefaultFullHttpResponse(HTTP_1_1, OK, Unpooled.wrappedBuffer(body)); + response.headers().set(CONTENT_LENGTH, body.length); + responseContext.write(response); + responseContext.flush(); + }); + } + + private void writeEarlyHintsWithTerminatorThenResponse() { + writeInformationalResponseWithTerminatorThenResponse(EARLY_HINTS); + } + + private void writeInformationalResponseWithTerminatorThenResponse(HttpResponseStatus status) { + ChannelHandlerContext responseContext = context; + if (responseContext == null) { + throw new IllegalStateException("no request context captured"); + } + String path; + synchronized (paths) { + path = paths.get(0); + } + responseContext + .executor() + .execute( + () -> { + responseContext.write(new DefaultHttpResponse(HTTP_1_1, status)); + responseContext.write(new DefaultLastHttpContent(Unpooled.EMPTY_BUFFER)); + + byte[] body = ("response " + path).getBytes(UTF_8); + DefaultFullHttpResponse response = + new DefaultFullHttpResponse(HTTP_1_1, OK, Unpooled.wrappedBuffer(body)); + response.headers().set(CONTENT_LENGTH, body.length); + responseContext.write(response); + responseContext.flush(); + }); + } + } +} diff --git a/dd-java-agent/instrumentation/netty/netty-common/src/main/java/datadog/trace/instrumentation/netty41/AttributeKeys.java b/dd-java-agent/instrumentation/netty/netty-common/src/main/java/datadog/trace/instrumentation/netty41/AttributeKeys.java index 3ad2b1f9e14..addf40b55d6 100644 --- a/dd-java-agent/instrumentation/netty/netty-common/src/main/java/datadog/trace/instrumentation/netty41/AttributeKeys.java +++ b/dd-java-agent/instrumentation/netty/netty-common/src/main/java/datadog/trace/instrumentation/netty41/AttributeKeys.java @@ -7,7 +7,6 @@ import datadog.trace.api.GenericClassValue; import datadog.trace.bootstrap.instrumentation.api.AgentSpan; import datadog.trace.bootstrap.instrumentation.websocket.HandlerContext; -import io.netty.handler.codec.http.HttpHeaders; import io.netty.util.AttributeKey; import java.util.concurrent.ConcurrentHashMap; import java.util.concurrent.ConcurrentMap; @@ -32,15 +31,6 @@ public final class AttributeKeys { public static final AttributeKey PARENT_CONTEXT_ATTRIBUTE_KEY = attributeKey("datadog.server.parent-context"); - public static final AttributeKey REQUEST_HEADERS_ATTRIBUTE_KEY = - attributeKey("datadog.server.request.headers"); - - public static final AttributeKey ANALYZED_RESPONSE_KEY = - attributeKey("datadog.server.analyzed_response"); - - public static final AttributeKey BLOCKED_RESPONSE_KEY = - attributeKey("datadog.server.blocked_response"); - public static final AttributeKey WEBSOCKET_SENDER_HANDLER_CONTEXT = attributeKey("datadog.server.websocket.sender.handler_context"); @@ -55,7 +45,7 @@ public final class AttributeKeys { * cassandra driver. */ @SuppressWarnings("unchecked") - private static AttributeKey attributeKey(final String key) { + static AttributeKey attributeKey(final String key) { ConcurrentMap> map = MAPS.get(AttributeKey.class); AttributeKey attributeKey = (AttributeKey) map.get(key); if (null == attributeKey) { diff --git a/dd-java-agent/instrumentation/netty/netty-common/src/main/java/datadog/trace/instrumentation/netty41/ServerRequestContext.java b/dd-java-agent/instrumentation/netty/netty-common/src/main/java/datadog/trace/instrumentation/netty41/ServerRequestContext.java new file mode 100644 index 00000000000..1d504e9aced --- /dev/null +++ b/dd-java-agent/instrumentation/netty/netty-common/src/main/java/datadog/trace/instrumentation/netty41/ServerRequestContext.java @@ -0,0 +1,242 @@ +package datadog.trace.instrumentation.netty41; + +import static datadog.trace.instrumentation.netty41.AttributeKeys.CONTEXT_ATTRIBUTE_KEY; + +import datadog.context.Context; +import datadog.trace.bootstrap.instrumentation.api.AgentSpan; +import io.netty.util.AttributeKey; +import io.netty.util.AttributeMap; +import java.util.ArrayDeque; +import java.util.Deque; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +/** Per-request server state stored on the channel until the matching response is written. */ +public final class ServerRequestContext { + /** + * Returns whether a new server request can be tracked on this channel (and may disable server + * tracing for the channel if the pending-context limit is exceeded). + */ + public static boolean canTrackRequest(final AttributeMap attributes) { + final Deque contexts = + attributes.attr(SERVER_REQUEST_CONTEXTS_ATTRIBUTE_KEY).get(); + return contexts == null || canAdd(attributes, contexts); + } + + /** Adds a request context to the queue tail. */ + public static ServerRequestContext add( + final AttributeMap attributes, final Context context, final String acceptHeader) { + return add(attributes, context, acceptHeader, false); + } + + /** Adds a request context to the queue tail. */ + public static ServerRequestContext add( + final AttributeMap attributes, + final Context context, + final String acceptHeader, + final boolean headRequest) { + final Deque contexts = getOrCreate(attributes); + if (!canAdd(attributes, contexts)) { + return null; + } + final ServerRequestContext serverContext = + new ServerRequestContext(context, acceptHeader, headRequest); + contexts.addLast(serverContext); + // The deque is authoritative for server request/response matching. CONTEXT_ATTRIBUTE_KEY is a + // context mirror of the current inbound request used by + // NettyChannelHandlerContextInstrumentation.FireAdvice and copied to HTTP/2 stream channels by + // Http2MultiplexHandlerStreamChannelInstrumentation. + attributes.attr(CONTEXT_ATTRIBUTE_KEY).set(context); + return serverContext; + } + + /** Returns the server request context for the next response. */ + public static ServerRequestContext nextResponse(final AttributeMap attributes) { + final Deque contexts = + attributes.attr(SERVER_REQUEST_CONTEXTS_ATTRIBUTE_KEY).get(); + // HTTP/1.1 responses are written in request order, including when requests are pipelined on one + // connection. + return contexts == null || isPoisoned(contexts) ? null : contexts.peekFirst(); + } + + /** Returns the server request context for the current inbound request. */ + public static ServerRequestContext currentRequest(final AttributeMap attributes) { + final Deque contexts = + attributes.attr(SERVER_REQUEST_CONTEXTS_ATTRIBUTE_KEY).get(); + return contexts == null || isPoisoned(contexts) ? null : contexts.peekLast(); + } + + /** Returns whether the channel is closing after an AppSec response block. */ + public static boolean isResponseBlocked(final AttributeMap attributes) { + return attributes.attr(BLOCKED_RESPONSE_ATTRIBUTE_KEY).get() != null; + } + + /** Marks the channel as closing after an AppSec response block. */ + public static void markResponseBlocked(final AttributeMap attributes) { + attributes.attr(BLOCKED_RESPONSE_ATTRIBUTE_KEY).set(Boolean.TRUE); + } + + /** Removes a completed or failed request context. */ + public static void remove( + final AttributeMap attributes, final ServerRequestContext serverContext) { + if (serverContext == null) { + return; + } + final Deque contexts = + attributes.attr(SERVER_REQUEST_CONTEXTS_ATTRIBUTE_KEY).get(); + if (contexts != null) { + if (isPoisoned(contexts)) { + return; + } + if (contexts.peekFirst() == serverContext) { + // Response completion consumes the queue head. + contexts.pollFirst(); + } else { + // Request-side failures normally remove the tail. Remove by value to cover later cleanup + // after additional pipelined requests were queued. + contexts.remove(serverContext); + } + final ServerRequestContext currentContext = contexts.peekLast(); + if (currentContext == null) { + attributes.attr(CONTEXT_ATTRIBUTE_KEY).remove(); + // No request-matching state remains. Drop the empty per-channel queue so idle + // keep-alive or upgraded channels do not retain it after the last HTTP request; + // getOrCreate will recreate it lazily if another request arrives. + attributes.attr(SERVER_REQUEST_CONTEXTS_ATTRIBUTE_KEY).remove(); + } else { + // Keep the context mirror pointed at the current inbound request after removing an older + // response context. It may still be copied to a new HTTP/2 stream channel. + attributes.attr(CONTEXT_ATTRIBUTE_KEY).set(currentContext.tracingContext()); + } + } + } + + /** Closes all pending request contexts on channel close. */ + public static void closeAll(final AttributeMap attributes) { + // The context mirror must not outlive the authoritative request queue. + attributes.attr(CONTEXT_ATTRIBUTE_KEY).remove(); + attributes.attr(BLOCKED_RESPONSE_ATTRIBUTE_KEY).remove(); + close(attributes.attr(SERVER_REQUEST_CONTEXTS_ATTRIBUTE_KEY).getAndRemove()); + } + + private static final int PIPELINING_LIMIT = 1000; + + private static final Logger log = LoggerFactory.getLogger(ServerRequestContext.class); + + /** Pending server request contexts for a channel. */ + private static final AttributeKey> + SERVER_REQUEST_CONTEXTS_ATTRIBUTE_KEY = + AttributeKeys.attributeKey("datadog.server.request.contexts"); + + private static final AttributeKey BLOCKED_RESPONSE_ATTRIBUTE_KEY = + AttributeKeys.attributeKey("datadog.server.blocked_response"); + + private static final Deque POISONED_CONTEXTS = new ArrayDeque<>(0); + + /** Creates the per-channel server request context queue. */ + private static Deque getOrCreate(final AttributeMap attributes) { + Deque contexts = + attributes.attr(SERVER_REQUEST_CONTEXTS_ATTRIBUTE_KEY).get(); + if (contexts == null) { + // Netty serializes handler callbacks for a channel on its EventLoop, so this queue does not + // need the allocation and atomic overhead of a concurrent collection. + contexts = new ArrayDeque<>(); + attributes.attr(SERVER_REQUEST_CONTEXTS_ATTRIBUTE_KEY).set(contexts); + } + return contexts; + } + + private static void close(final Deque contexts) { + if (contexts == null || isPoisoned(contexts)) { + return; + } + ServerRequestContext context; + while ((context = contexts.pollFirst()) != null) { + try { + final AgentSpan span = AgentSpan.fromContext(context.tracingContext()); + if (span != null && span.phasedFinish()) { + // These contexts no longer have a response handler path that can finish the span. + span.publish(); + } + } catch (final Throwable ignored) { + } + } + } + + private static boolean canAdd( + final AttributeMap attributes, final Deque contexts) { + if (isPoisoned(contexts)) { + return false; + } + // If this limit is exceeded, stop tracing on the channel and drain the deque. This suggests + // contexts are not being removed, for example because the server stopped writing responses. + if (contexts.size() >= PIPELINING_LIMIT) { + final int pendingContexts = contexts.size(); + close(contexts); + attributes.attr(SERVER_REQUEST_CONTEXTS_ATTRIBUTE_KEY).set(POISONED_CONTEXTS); + attributes.attr(CONTEXT_ATTRIBUTE_KEY).remove(); + log.error( + "Too many pending Netty server request contexts on a channel; " + + "closing {} contexts and disabling Netty server tracing on that channel " + + "(limit: {})", + pendingContexts, + PIPELINING_LIMIT); + return false; + } + return true; + } + + private static boolean isPoisoned(final Deque contexts) { + return contexts == POISONED_CONTEXTS; + } + + private final Context tracingContext; + private final String acceptHeader; + private final boolean headRequest; + private boolean responseStarted; + private boolean responseAnalyzed; + private Object deferredBlockResponse; + + public Context tracingContext() { + return tracingContext; + } + + public String acceptHeader() { + return acceptHeader; + } + + public boolean isHeadRequest() { + return headRequest; + } + + public boolean isResponseStarted() { + return responseStarted; + } + + public void markResponseStarted() { + responseStarted = true; + } + + public boolean isResponseAnalyzed() { + return responseAnalyzed; + } + + public void markResponseAnalyzed() { + responseAnalyzed = true; + } + + public Object deferredBlockResponse() { + return deferredBlockResponse; + } + + public void deferBlockResponse(final Object deferredBlockResponse) { + this.deferredBlockResponse = deferredBlockResponse; + } + + private ServerRequestContext( + final Context tracingContext, final String acceptHeader, final boolean headRequest) { + this.tracingContext = tracingContext; + this.acceptHeader = acceptHeader; + this.headRequest = headRequest; + } +}