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
132 changes: 73 additions & 59 deletions xds/src/main/java/io/grpc/xds/ExternalProcessorClientInterceptor.java

Large diffs are not rendered by default.

14 changes: 14 additions & 0 deletions xds/src/main/java/io/grpc/xds/ExternalProcessorFilter.java
Original file line number Diff line number Diff line change
Expand Up @@ -134,6 +134,20 @@ public ClientInterceptor buildClientInterceptor(FilterConfig filterConfig,
extProcFilterConfig, cachedChannelManager, scheduler, context);
}

@Override
public boolean requiresPayloadAccess(
FilterConfig filterConfig, @Nullable FilterConfig overrideConfig) {
ExternalProcessorFilterConfig extProcFilterConfig =
(ExternalProcessorFilterConfig) filterConfig;
if (overrideConfig != null) {
extProcFilterConfig = mergeConfigs(extProcFilterConfig,
(ExternalProcessorFilterOverrideConfig) overrideConfig);
}
ProcessingMode mode = extProcFilterConfig.getExternalProcessor().getProcessingMode();
return mode.getRequestBodyMode() != ProcessingMode.BodySendMode.NONE
|| mode.getResponseBodyMode() != ProcessingMode.BodySendMode.NONE;
}

private static ExternalProcessorFilterConfig mergeConfigs(
ExternalProcessorFilterConfig extProcFilterConfig,
ExternalProcessorFilterOverrideConfig extProcFilterConfigOverride) {
Expand Down
12 changes: 12 additions & 0 deletions xds/src/main/java/io/grpc/xds/Filter.java
Original file line number Diff line number Diff line change
Expand Up @@ -124,6 +124,18 @@ default ServerInterceptor buildServerInterceptor(
return null;
}

/**
* Returns true if this filter requires access to the request or response message payloads for
* the given configuration.
*
* <p>When true, interceptors that provide raw message payload access (e.g. {@code
* RawMessageClientInterceptor}) will be installed in the client interceptor chain.
*/
default boolean requiresPayloadAccess(
FilterConfig config, @Nullable FilterConfig overrideConfig) {
return false;
}

/**
* Releases filter resources like shared resources and remote connections.
*
Expand Down
44 changes: 41 additions & 3 deletions xds/src/main/java/io/grpc/xds/XdsNameResolver.java
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@
import com.google.common.collect.ImmutableMap;
import com.google.common.collect.Sets;
import com.google.gson.Gson;
import com.google.protobuf.ByteString;
import com.google.protobuf.util.Durations;
import io.grpc.Attributes;
import io.grpc.CallOptions;
Expand All @@ -35,6 +36,7 @@
import io.grpc.ClientCall;
import io.grpc.ClientInterceptor;
import io.grpc.ClientInterceptors;
import io.grpc.Drainable;
import io.grpc.ForwardingClientCall.SimpleForwardingClientCall;
import io.grpc.ForwardingClientCallListener.SimpleForwardingClientCallListener;
import io.grpc.InternalConfigSelector;
Expand Down Expand Up @@ -69,6 +71,8 @@
import io.grpc.xds.client.XdsInitializationException;
import io.grpc.xds.client.XdsLogger;
import io.grpc.xds.client.XdsLogger.XdsLogLevel;
import io.grpc.xds.internal.extproc.KnownLengthInputStream;
import java.io.IOException;
import java.io.InputStream;
import java.util.ArrayList;
import java.util.Collections;
Expand Down Expand Up @@ -889,6 +893,7 @@ private ClientInterceptor createFilters(
selectedOverrideConfigs.putAll(weightedCluster.filterConfigOverrides());
}

boolean anyFilterRequiresPayload = false;
ImmutableList.Builder<ClientInterceptor> filterInterceptors = ImmutableList.builder();
for (NamedFilterConfig namedFilter : filterConfigs) {
String name = namedFilter.name;
Expand All @@ -903,12 +908,14 @@ private ClientInterceptor createFilters(

if (interceptor != null) {
filterInterceptors.add(interceptor);
if (filter.requiresPayloadAccess(config, overrideConfig)) {
anyFilterRequiresPayload = true;
}
}
}

ImmutableList.Builder<ClientInterceptor> withRawMessage = ImmutableList.builder();
if (GrpcUtil.getFlag("GRPC_EXPERIMENTAL_XDS_EXT_PROC_ON_CLIENT", false)
|| GrpcUtil.getFlag("GRPC_EXPERIMENTAL_XDS_EXT_PROC_ON_SERVER", false)) {
if (anyFilterRequiresPayload) {
withRawMessage.add(new RawMessageClientInterceptor());
}
withRawMessage.addAll(filterInterceptors.build());
Expand Down Expand Up @@ -1134,6 +1141,13 @@ static final class RawMessageClientInterceptor implements ClientInterceptor {
new MethodDescriptor.Marshaller<InputStream>() {
@Override
public InputStream stream(InputStream value) {
// For retry attempts, RetriableStream calls stream(value) once per attempt.
// Returning a fresh KnownLengthInputStream wrapping the immutable ByteString ensures
// each retry attempt reads from the beginning of the payload rather than an already
// drained stream.
if (value instanceof KnownLengthInputStream) {
return new KnownLengthInputStream(((KnownLengthInputStream) value).getByteString());
}
return value;
}

Expand Down Expand Up @@ -1192,7 +1206,31 @@ public void halfClose() {

@Override
public void sendMessage(ReqT message) {
rawCall.sendMessage(method.getRequestMarshaller().stream(message));
InputStream stream = method.getRequestMarshaller().stream(message);
ByteString byteString;
try {
if (stream instanceof Drainable) {
int size = stream.available();
ByteString.Output output =
size > 0 ? ByteString.newOutput(size) : ByteString.newOutput();
((Drainable) stream).drainTo(output);
byteString = output.toByteString();
} else {
byteString = ByteString.readFrom(stream);
}
} catch (IOException e) {
throw Status.INTERNAL
.withDescription("Failed to read message for raw message interceptor")
.withCause(e)
.asRuntimeException();
} finally {
try {
stream.close();
} catch (IOException ignored) {
// ignore
}
}
rawCall.sendMessage(new KnownLengthInputStream(byteString));
}

@Override
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -25,12 +25,18 @@
* An {@link InputStream} backed by a {@link ByteString} that implements {@link KnownLength}.
*/
public final class KnownLengthInputStream extends InputStream implements KnownLength {
private final ByteString byteString;
private final InputStream delegate;

public KnownLengthInputStream(ByteString byteString) {
this.byteString = byteString;
this.delegate = byteString.newInput();
}

public ByteString getByteString() {
return byteString;
}

@Override
public int read() throws IOException {
return delegate.read();
Expand Down
112 changes: 112 additions & 0 deletions xds/src/test/java/io/grpc/xds/ExternalProcessorFilterTest.java
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@
import io.envoyproxy.envoy.extensions.filters.http.ext_proc.v3.ExtProcPerRoute;
import io.envoyproxy.envoy.extensions.filters.http.ext_proc.v3.ExternalProcessor;
import io.envoyproxy.envoy.extensions.filters.http.ext_proc.v3.ProcessingMode;
import io.grpc.MetricRecorder;
import io.grpc.NameResolver;
import io.grpc.NameResolverProvider;
import io.grpc.NameResolverRegistry;
Expand Down Expand Up @@ -341,4 +342,115 @@ public void givenInvalidProto_whenParseFilterConfigOverride_thenReturnsError() t

assertThat(result.errorDetail).contains("Invalid proto:");
}

@Test
public void requiresPayloadAccess_defaultProcessingMode_returnsFalse() throws Exception {
ExternalProcessor proto = createBaseProto(extProcServerName).build();
ExternalProcessorFilterConfig config =
provider.parseFilterConfig(Any.pack(proto), filterContext).config;

try (ExternalProcessorFilter filter = provider.newInstance(
Filter.FilterContext.create("test-filter", new MetricRecorder() {}))) {
assertThat(filter.requiresPayloadAccess(config, null)).isFalse();
}
}

@Test
public void requiresPayloadAccess_explicitNone_returnsFalse() throws Exception {
ExternalProcessor proto = createBaseProto(extProcServerName)
.setProcessingMode(ProcessingMode.newBuilder()
.setRequestBodyMode(ProcessingMode.BodySendMode.NONE)
.setResponseBodyMode(ProcessingMode.BodySendMode.NONE)
.build())
.build();
ExternalProcessorFilterConfig config =
provider.parseFilterConfig(Any.pack(proto), filterContext).config;

try (ExternalProcessorFilter filter = provider.newInstance(
Filter.FilterContext.create("test-filter", new MetricRecorder() {}))) {
assertThat(filter.requiresPayloadAccess(config, null)).isFalse();
}
}

@Test
public void requiresPayloadAccess_requestBodyGrpc_returnsTrue() throws Exception {
ExternalProcessor proto = createBaseProto(extProcServerName)
.setProcessingMode(ProcessingMode.newBuilder()
.setRequestBodyMode(ProcessingMode.BodySendMode.GRPC)
.build())
.build();
ExternalProcessorFilterConfig config =
provider.parseFilterConfig(Any.pack(proto), filterContext).config;

try (ExternalProcessorFilter filter = provider.newInstance(
Filter.FilterContext.create("test-filter", new MetricRecorder() {}))) {
assertThat(filter.requiresPayloadAccess(config, null)).isTrue();
}
}

@Test
public void requiresPayloadAccess_responseBodyGrpc_returnsTrue() throws Exception {
ExternalProcessor proto = createBaseProto(extProcServerName)
.setProcessingMode(ProcessingMode.newBuilder()
.setResponseBodyMode(ProcessingMode.BodySendMode.GRPC)
.setResponseTrailerMode(ProcessingMode.HeaderSendMode.SEND)
.build())
.build();
ExternalProcessorFilterConfig config =
provider.parseFilterConfig(Any.pack(proto), filterContext).config;

try (ExternalProcessorFilter filter = provider.newInstance(
Filter.FilterContext.create("test-filter", new MetricRecorder() {}))) {
assertThat(filter.requiresPayloadAccess(config, null)).isTrue();
}
}

@Test
public void requiresPayloadAccess_overrideTurnsOnPayloadAccess_returnsTrue() throws Exception {
ExternalProcessor proto = createBaseProto(extProcServerName).build();
ExternalProcessorFilterConfig config =
provider.parseFilterConfig(Any.pack(proto), filterContext).config;

ExtProcPerRoute perRoute = ExtProcPerRoute.newBuilder()
.setOverrides(ExtProcOverrides.newBuilder()
.setProcessingMode(ProcessingMode.newBuilder()
.setRequestBodyMode(ProcessingMode.BodySendMode.GRPC)
.build())
.build())
.build();
ExternalProcessorFilterOverrideConfig overrideConfig =
provider.parseFilterConfigOverride(Any.pack(perRoute), filterContext).config;

try (ExternalProcessorFilter filter = provider.newInstance(
Filter.FilterContext.create("test-filter", new MetricRecorder() {}))) {
assertThat(filter.requiresPayloadAccess(config, overrideConfig)).isTrue();
}
}

@Test
public void requiresPayloadAccess_overrideTurnsOffPayloadAccess_returnsFalse() throws Exception {
ExternalProcessor proto = createBaseProto(extProcServerName)
.setProcessingMode(ProcessingMode.newBuilder()
.setRequestBodyMode(ProcessingMode.BodySendMode.GRPC)
.build())
.build();
ExternalProcessorFilterConfig config =
provider.parseFilterConfig(Any.pack(proto), filterContext).config;

ExtProcPerRoute perRoute = ExtProcPerRoute.newBuilder()
.setOverrides(ExtProcOverrides.newBuilder()
.setProcessingMode(ProcessingMode.newBuilder()
.setRequestBodyMode(ProcessingMode.BodySendMode.NONE)
.setResponseBodyMode(ProcessingMode.BodySendMode.NONE)
.build())
.build())
.build();
ExternalProcessorFilterOverrideConfig overrideConfig =
provider.parseFilterConfigOverride(Any.pack(perRoute), filterContext).config;

try (ExternalProcessorFilter filter = provider.newInstance(
Filter.FilterContext.create("test-filter", new MetricRecorder() {}))) {
assertThat(filter.requiresPayloadAccess(config, overrideConfig)).isFalse();
}
}
}
47 changes: 45 additions & 2 deletions xds/src/test/java/io/grpc/xds/StatefulFilter.java
Original file line number Diff line number Diff line change
Expand Up @@ -21,10 +21,16 @@

import com.google.common.collect.ImmutableList;
import com.google.protobuf.Message;
import io.grpc.CallOptions;
import io.grpc.Channel;
import io.grpc.ClientCall;
import io.grpc.ClientInterceptor;
import io.grpc.MethodDescriptor;
import io.grpc.ServerInterceptor;
import java.util.ConcurrentModificationException;
import java.util.concurrent.ConcurrentHashMap;
import java.util.concurrent.ConcurrentMap;
import java.util.concurrent.ScheduledExecutorService;
import java.util.concurrent.atomic.AtomicBoolean;
import java.util.stream.IntStream;
import javax.annotation.Nullable;
Expand All @@ -38,10 +44,22 @@ class StatefulFilter implements Filter {
private final AtomicBoolean shutdown = new AtomicBoolean();

final int idx;
private final boolean requiresPayloadAccess;
@Nullable volatile String lastCfg = null;

public StatefulFilter(int idx) {
this(idx, false);
}

public StatefulFilter(int idx, boolean requiresPayloadAccess) {
this.idx = idx;
this.requiresPayloadAccess = requiresPayloadAccess;
}

@Override
public boolean requiresPayloadAccess(
FilterConfig config, @Nullable FilterConfig overrideConfig) {
return requiresPayloadAccess;
}

public boolean isShutdown() {
Expand All @@ -67,6 +85,21 @@ public ServerInterceptor buildServerInterceptor(
return null;
}

@Nullable
@Override
public ClientInterceptor buildClientInterceptor(
FilterConfig config,
@Nullable FilterConfig overrideConfig,
ScheduledExecutorService scheduler) {
return new ClientInterceptor() {
@Override
public <ReqT, RespT> ClientCall<ReqT, RespT> interceptCall(
MethodDescriptor<ReqT, RespT> method, CallOptions callOptions, Channel next) {
return next.newCall(method, callOptions);
}
};
}

@Override
public String toString() {
StringBuilder sb = new StringBuilder().append("StatefulFilter{")
Expand All @@ -80,16 +113,26 @@ public String toString() {
static final class Provider implements Filter.Provider {

private final String typeUrl;
private final boolean requiresPayloadAccess;
private final ConcurrentMap<Integer, StatefulFilter> instances = new ConcurrentHashMap<>();

volatile int counter;

Provider() {
this(DEFAULT_TYPE_URL);
this(DEFAULT_TYPE_URL, false);
}

Provider(boolean requiresPayloadAccess) {
this(DEFAULT_TYPE_URL, requiresPayloadAccess);
}

Provider(String typeUrl) {
this(typeUrl, false);
}

Provider(String typeUrl, boolean requiresPayloadAccess) {
this.typeUrl = typeUrl;
this.requiresPayloadAccess = requiresPayloadAccess;
}

@Override
Expand All @@ -109,7 +152,7 @@ public boolean isServerFilter() {

@Override
public synchronized StatefulFilter newInstance(FilterContext context) {
StatefulFilter filter = new StatefulFilter(counter++);
StatefulFilter filter = new StatefulFilter(counter++, requiresPayloadAccess);
instances.put(filter.idx, filter);
return filter;
}
Expand Down
Loading
Loading