Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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);
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,7 @@ public ElementMatcher<TypeDescription> hierarchyMatcher() {
public String[] helperClassNames() {
return new String[] {
packageName + ".AttributeKeys",
packageName + ".ServerRequestContext",
// client helpers
packageName + ".client.NettyHttpClientDecorator",
packageName + ".client.NettyResponseInjectAdapter",
Expand All @@ -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"
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -47,12 +46,14 @@ public ElementMatcher<TypeDescription> 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"
Expand All @@ -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;
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -72,6 +72,7 @@ public ElementMatcher<TypeDescription> hierarchyMatcher() {
public String[] helperClassNames() {
return new String[] {
packageName + ".AttributeKeys",
packageName + ".ServerRequestContext",
// client helpers
packageName + ".client.NettyHttpClientDecorator",
packageName + ".client.NettyResponseInjectAdapter",
Expand All @@ -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",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -15,20 +16,25 @@
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;
import org.slf4j.LoggerFactory;

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;
private final int statusCode;
private final BlockingContentType bct;
private final Map<String, String> extraHeaders;
private final String securityResponseId;
private final ServerRequestContext serverContext;

private boolean hasBlockedAlready;

Expand All @@ -37,21 +43,26 @@ public BlockingResponseHandler(
int statusCode,
BlockingContentType bct,
Map<String, String> 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
Expand All @@ -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<String, String> 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<String, String> 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<String, String> 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<String, String> 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();
Expand Down
Loading