diff --git a/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/web/client/TraceWebClientAutoConfiguration.java b/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/web/client/TraceWebClientAutoConfiguration.java index 556d57389..4fa2b851f 100644 --- a/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/web/client/TraceWebClientAutoConfiguration.java +++ b/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/web/client/TraceWebClientAutoConfiguration.java @@ -17,9 +17,14 @@ package org.springframework.cloud.sleuth.instrument.web.client; import java.io.IOException; +import java.nio.charset.Charset; +import java.nio.file.Path; import java.util.ArrayList; import java.util.List; +import java.util.Map; +import java.util.Set; import java.util.concurrent.atomic.AtomicReference; +import java.util.function.Consumer; import java.util.function.Function; import brave.Span; @@ -31,8 +36,12 @@ import brave.httpclient.TracingHttpClientBuilder; import brave.propagation.Propagation; import brave.propagation.TraceContext; import brave.spring.web.TracingClientHttpRequestInterceptor; +import io.netty.buffer.ByteBuf; +import io.netty.buffer.ByteBufAllocator; import io.netty.handler.codec.http.HttpHeaders; import io.netty.handler.codec.http.HttpMethod; +import io.netty.handler.codec.http.HttpVersion; +import io.netty.handler.codec.http.cookie.Cookie; import org.apache.commons.logging.Log; import org.apache.commons.logging.LogFactory; import org.apache.http.impl.client.HttpClientBuilder; @@ -42,6 +51,7 @@ import org.aspectj.lang.annotation.Around; import org.aspectj.lang.annotation.Aspect; import org.aspectj.lang.annotation.Pointcut; import org.reactivestreams.Publisher; +import org.reactivestreams.Subscriber; import org.springframework.beans.BeansException; import org.springframework.beans.factory.BeanFactory; import org.springframework.beans.factory.ListableBeanFactory; @@ -66,10 +76,16 @@ import org.springframework.http.client.ClientHttpResponse; import org.springframework.security.oauth2.client.OAuth2RestTemplate; import org.springframework.web.client.RestTemplate; import org.springframework.web.reactive.function.client.WebClient; +import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; +import reactor.ipc.netty.NettyContext; +import reactor.ipc.netty.NettyOutbound; +import reactor.ipc.netty.NettyPipeline; +import reactor.ipc.netty.channel.data.FileChunkedStrategy; import reactor.ipc.netty.http.client.HttpClient; import reactor.ipc.netty.http.client.HttpClientRequest; import reactor.ipc.netty.http.client.HttpClientResponse; +import reactor.ipc.netty.http.websocket.WebsocketOutbound; /** * {@link org.springframework.boot.autoconfigure.EnableAutoConfiguration @@ -362,14 +378,18 @@ class TracingHttpClientInstrumentation { Function> combinedFunction = req -> { try (Tracer.SpanInScope spanInScope = this.tracer.withSpanInScope(currentSpan)) { - io.netty.handler.codec.http.HttpHeaders headers = req + io.netty.handler.codec.http.HttpHeaders originalHeaders = req + .requestHeaders().copy(); + io.netty.handler.codec.http.HttpHeaders tracedHeaders = req .requestHeaders(); - span.set(this.handler.handleSend(this.injector, headers, req)); + span.set(this.handler.handleSend(this.injector, tracedHeaders, req)); + io.netty.handler.codec.http.HttpHeaders addedHeaders = tracedHeaders.copy(); + originalHeaders.forEach(header -> addedHeaders.remove(header.getKey())); try (Tracer.SpanInScope clientInScope = this.tracer.withSpanInScope(span.get())) { if (log.isDebugEnabled()) { log.debug("Created a new client span for Netty client"); } - return handle(handler, req); + return handle(handler, new TracedHttpClientRequest(req, addedHeaders)); } } }; @@ -388,6 +408,227 @@ class TracingHttpClientInstrumentation { }); } + /** + * The `org.springframework.cloud.gateway.filter.NettyRoutingFilter` in SC Gateway + * is adding only these headers that were set when the request came in. That means + * that adding any additional headers (via instrumentation) is completely ignored. + * That's why we're wrapping the `HttpClientRequest` in such a wrapper that + * when `setHeaders` is called (that clears any current headers), will also add + * the tracing headers + */ + static class TracedHttpClientRequest implements HttpClientRequest { + private HttpClientRequest delegate; + private final io.netty.handler.codec.http.HttpHeaders addedHeaders; + + TracedHttpClientRequest(HttpClientRequest delegate, HttpHeaders addedHeaders) { + this.delegate = delegate; + this.addedHeaders = addedHeaders; + } + + @Override public HttpClientRequest addCookie(Cookie cookie) { + this.delegate = this.delegate.addCookie(cookie); + return this; + } + + @Override public HttpClientRequest addHeader(CharSequence name, + CharSequence value) { + this.delegate = this.delegate.addHeader(name, value); + return this; + } + + @Override public HttpClientRequest context( + Consumer contextCallback) { + this.delegate = this.delegate.context(contextCallback); + return this; + } + + @Override public HttpClientRequest chunkedTransfer(boolean chunked) { + this.delegate = this.delegate.chunkedTransfer(chunked); + return this; + } + + @Override public HttpClientRequest options( + Consumer configurator) { + this.delegate = this.delegate.options(configurator); + return this; + } + + @Override public HttpClientRequest followRedirect() { + this.delegate = this.delegate.followRedirect(); + return this; + } + + @Override public HttpClientRequest failOnClientError(boolean shouldFail) { + this.delegate = this.delegate.failOnClientError(shouldFail); + return this; + } + + @Override public HttpClientRequest failOnServerError(boolean shouldFail) { + this.delegate = this.delegate.failOnServerError(shouldFail); + return this; + } + + @Override public boolean hasSentHeaders() { + return this.delegate.hasSentHeaders(); + } + + @Override public HttpClientRequest header(CharSequence name, CharSequence value) { + this.delegate = this.delegate.header(name, value); + return this; + } + + @Override public HttpClientRequest headers(HttpHeaders headers) { + HttpHeaders copy = headers.copy(); + copy.add(this.addedHeaders); + this.delegate = this.delegate.headers(copy); + return this; + } + + @Override public boolean isFollowRedirect() { + return this.delegate.isFollowRedirect(); + } + + @Override public HttpClientRequest keepAlive(boolean keepAlive) { + this.delegate = this.delegate.keepAlive(keepAlive); + return this; + } + + @Override public HttpClientRequest onWriteIdle(long idleTimeout, + Runnable onWriteIdle) { + this.delegate = this.delegate.onWriteIdle(idleTimeout, onWriteIdle); + return this; + } + + @Override public String[] redirectedFrom() { + return this.delegate.redirectedFrom(); + } + + @Override public HttpHeaders requestHeaders() { + return this.delegate.requestHeaders(); + } + + @Override public Mono send() { + return this.delegate.send(); + } + + @Override public Flux sendForm(Consumer
formCallback) { + return this.delegate.sendForm(formCallback); + } + + @Override public NettyOutbound sendHeaders() { + return this.delegate.sendHeaders(); + } + + @Override public WebsocketOutbound sendWebsocket() { + return this.delegate.sendWebsocket(); + } + + @Override public WebsocketOutbound sendWebsocket(String subprotocols) { + return this.delegate.sendWebsocket(subprotocols); + } + + @Override public ByteBufAllocator alloc() { + return this.delegate.alloc(); + } + + @Override public NettyContext context() { + return this.delegate.context(); + } + + @Override public FileChunkedStrategy getFileChunkedStrategy() { + return this.delegate.getFileChunkedStrategy(); + } + + @Override public Mono neverComplete() { + return this.delegate.neverComplete(); + } + + @Override public NettyOutbound send(Publisher dataStream) { + return this.delegate.send(dataStream); + } + + @Override public NettyOutbound sendByteArray( + Publisher dataStream) { + return this.delegate.sendByteArray(dataStream); + } + + @Override public NettyOutbound sendFile(Path file) { + return this.delegate.sendFile(file); + } + + @Override public NettyOutbound sendFile(Path file, long position, long count) { + return this.delegate.sendFile(file, position, count); + } + + @Override public NettyOutbound sendFileChunked(Path file, long position, + long count) { + return this.delegate.sendFileChunked(file, position, count); + } + + @Override public NettyOutbound sendGroups( + Publisher> dataStreams) { + return this.delegate.sendGroups(dataStreams); + } + + @Override public NettyOutbound sendObject(Publisher dataStream) { + return this.delegate.sendObject(dataStream); + } + + @Override public NettyOutbound sendObject(Object msg) { + return this.delegate.sendObject(msg); + } + + @Override public NettyOutbound sendString( + Publisher dataStream) { + return this.delegate.sendString(dataStream); + } + + @Override public NettyOutbound sendString(Publisher dataStream, + Charset charset) { + return this.delegate.sendString(dataStream, charset); + } + + @Override public void subscribe(Subscriber s) { + this.delegate.subscribe(s); + } + + @Override public Mono then() { + return this.delegate.then(); + } + + @Override public NettyOutbound then(Publisher other) { + return this.delegate.then(other); + } + + @Override public Map> cookies() { + return this.delegate.cookies(); + } + + @Override public boolean isKeepAlive() { + return this.delegate.isKeepAlive(); + } + + @Override public boolean isWebsocket() { + return this.delegate.isWebsocket(); + } + + @Override public HttpMethod method() { + return this.delegate.method(); + } + + @Override public String path() { + return this.delegate.path(); + } + + @Override public String uri() { + return this.delegate.uri(); + } + + @Override public HttpVersion version() { + return this.delegate.version(); + } + } + private Publisher handle( Function> handler, HttpClientRequest req) {