diff --git a/api/src/main/java/io/grpc/Contexts.java b/api/src/main/java/io/grpc/Contexts.java index c62ffc80a38..9c3697dd6d4 100644 --- a/api/src/main/java/io/grpc/Contexts.java +++ b/api/src/main/java/io/grpc/Contexts.java @@ -118,6 +118,16 @@ public void onReady() { context.detach(previous); } } + + @Override + public void onEvent(Object event) { + Context previous = context.attach(); + try { + super.onEvent(event); + } finally { + context.detach(previous); + } + } } /** diff --git a/api/src/main/java/io/grpc/PartialForwardingServerCall.java b/api/src/main/java/io/grpc/PartialForwardingServerCall.java index a313407b23e..8c8f53cf93c 100644 --- a/api/src/main/java/io/grpc/PartialForwardingServerCall.java +++ b/api/src/main/java/io/grpc/PartialForwardingServerCall.java @@ -87,6 +87,11 @@ public SecurityLevel getSecurityLevel() { return delegate().getSecurityLevel(); } + @Override + public void triggerEvent(Object event) { + delegate().triggerEvent(event); + } + @Override public String toString() { return MoreObjects.toStringHelper(this).add("delegate", delegate()).toString(); diff --git a/api/src/main/java/io/grpc/PartialForwardingServerCallListener.java b/api/src/main/java/io/grpc/PartialForwardingServerCallListener.java index ca2fd0058c9..23e93bb065e 100644 --- a/api/src/main/java/io/grpc/PartialForwardingServerCallListener.java +++ b/api/src/main/java/io/grpc/PartialForwardingServerCallListener.java @@ -50,6 +50,11 @@ public void onReady() { delegate().onReady(); } + @Override + public void onEvent(Object event) { + delegate().onEvent(event); + } + @Override public String toString() { return MoreObjects.toStringHelper(this).add("delegate", delegate()).toString(); diff --git a/api/src/main/java/io/grpc/ServerCall.java b/api/src/main/java/io/grpc/ServerCall.java index 3db8ac30e83..2e6ee07a23f 100644 --- a/api/src/main/java/io/grpc/ServerCall.java +++ b/api/src/main/java/io/grpc/ServerCall.java @@ -100,6 +100,20 @@ public void onComplete() {} * another {@code onReady()} callback. */ public void onReady() {} + + /** + * A custom event has been triggered by the call. + * + *

This callback is guaranteed to run on the call's executor, serialized with other + * callbacks (like {@link #onMessage}, {@link #onHalfClose}). This means the implementation + * does not need internal synchronization to access call-specific state. + * + * @param event the triggered event. + */ + @ExperimentalApi("https://github.com/grpc/grpc-java/issues/12979") + public void onEvent(Object event) { + // Default no-op + } } /** @@ -262,6 +276,20 @@ public String getAuthority() { return null; } + /** + * Triggers a custom event to be processed by the listener. + * The event will be delivered to {@link Listener#onEvent(Object)} on the call's executor. + * + *

This method is thread-safe and can be called from any thread. No events will be delivered + * after the RPC is cancelled or completed. + * + * @param event the event to trigger. + */ + @ExperimentalApi("https://github.com/grpc/grpc-java/issues/12979") + public void triggerEvent(Object event) { + // Default no-op + } + /** * The {@link MethodDescriptor} for the call. */ diff --git a/api/src/test/java/io/grpc/ContextsTest.java b/api/src/test/java/io/grpc/ContextsTest.java index ec9dc3929a2..974b1aff0c2 100644 --- a/api/src/test/java/io/grpc/ContextsTest.java +++ b/api/src/test/java/io/grpc/ContextsTest.java @@ -82,6 +82,11 @@ public void interceptCall_basic() { assertSame(uniqueContext, Context.current()); methodCalls.add(5); } + + @Override public void onEvent(Object event) { + assertSame(uniqueContext, Context.current()); + methodCalls.add(6); + } }; ServerCall.Listener wrapped = interceptCall(uniqueContext, call, headers, new ServerCallHandler() { @@ -101,7 +106,8 @@ public ServerCall.Listener startCall( wrapped.onCancel(); wrapped.onComplete(); wrapped.onReady(); - assertEquals(Arrays.asList(1, 2, 3, 4, 5), methodCalls); + wrapped.onEvent(new Object()); + assertEquals(Arrays.asList(1, 2, 3, 4, 5, 6), methodCalls); assertSame(origContext, Context.current()); } @@ -145,6 +151,10 @@ public void interceptCall_restoresIfListenerThrows() { @Override public void onReady() { throw new RuntimeException(); } + + @Override public void onEvent(Object event) { + throw new RuntimeException(); + } }; ServerCall.Listener wrapped = interceptCall(uniqueContext, call, headers, new ServerCallHandler() { @@ -180,6 +190,11 @@ public ServerCall.Listener startCall( fail("Exception expected"); } catch (RuntimeException expected) { } + try { + wrapped.onEvent(new Object()); + fail("Exception expected"); + } catch (RuntimeException expected) { + } assertSame(origContext, Context.current()); } diff --git a/binder/src/main/java/io/grpc/binder/internal/Inbound.java b/binder/src/main/java/io/grpc/binder/internal/Inbound.java index 83fc8273d6f..83decf4a89a 100644 --- a/binder/src/main/java/io/grpc/binder/internal/Inbound.java +++ b/binder/src/main/java/io/grpc/binder/internal/Inbound.java @@ -668,6 +668,19 @@ protected void deliverCloseAbnormal(Status status) { listener.closed(status); } + void triggerEvent(Object event) { + ServerStreamListener localListener; + synchronized (this) { + if (isClosed()) { + return; + } + localListener = listener; + } + if (localListener != null) { + localListener.triggerEvent(event); + } + } + @GuardedBy("this") void onCloseSent(Status status) { if (!isClosed()) { diff --git a/binder/src/main/java/io/grpc/binder/internal/MultiMessageServerStream.java b/binder/src/main/java/io/grpc/binder/internal/MultiMessageServerStream.java index f54769caefa..7a57138ce22 100644 --- a/binder/src/main/java/io/grpc/binder/internal/MultiMessageServerStream.java +++ b/binder/src/main/java/io/grpc/binder/internal/MultiMessageServerStream.java @@ -175,6 +175,11 @@ public void setDecompressor(Decompressor decompressor) { // Ignore. } + @Override + public void triggerEvent(Object event) { + inbound.triggerEvent(event); + } + @Override public void optimizeForDirectExecutor() { // Ignore. diff --git a/binder/src/main/java/io/grpc/binder/internal/SingleMessageServerStream.java b/binder/src/main/java/io/grpc/binder/internal/SingleMessageServerStream.java index 383bd7a2593..5f1dd511f73 100644 --- a/binder/src/main/java/io/grpc/binder/internal/SingleMessageServerStream.java +++ b/binder/src/main/java/io/grpc/binder/internal/SingleMessageServerStream.java @@ -167,6 +167,11 @@ public void setDecompressor(Decompressor decompressor) { // Ignore. } + @Override + public void triggerEvent(Object event) { + inbound.triggerEvent(event); + } + @Override public void optimizeForDirectExecutor() { // Ignore. diff --git a/core/src/main/java/io/grpc/internal/AbstractServerStream.java b/core/src/main/java/io/grpc/internal/AbstractServerStream.java index c468cba978a..67dfdc93d42 100644 --- a/core/src/main/java/io/grpc/internal/AbstractServerStream.java +++ b/core/src/main/java/io/grpc/internal/AbstractServerStream.java @@ -173,6 +173,16 @@ public final void setListener(ServerStreamListener serverStreamListener) { transportState().setListener(serverStreamListener); } + @Override + public final void triggerEvent(final Object event) { + transportState().runOnTransportThread(new Runnable() { + @Override + public void run() { + transportState().triggerEvent(event); + } + }); + } + @Override public StatsTraceContext statsTraceContext() { return statsTraceCtx; @@ -259,6 +269,13 @@ public void deframerClosed(boolean hasPartialMessage) { + public final void triggerEvent(Object event) { + if (listenerClosed) { + return; + } + listener().triggerEvent(event); + } + @Override protected ServerStreamListener listener() { return listener; diff --git a/core/src/main/java/io/grpc/internal/ServerCallImpl.java b/core/src/main/java/io/grpc/internal/ServerCallImpl.java index e224384ce8f..6e894371be1 100644 --- a/core/src/main/java/io/grpc/internal/ServerCallImpl.java +++ b/core/src/main/java/io/grpc/internal/ServerCallImpl.java @@ -254,6 +254,11 @@ public MethodDescriptor getMethodDescriptor() { return method; } + @Override + public void triggerEvent(Object event) { + stream.triggerEvent(event); + } + @Override public SecurityLevel getSecurityLevel() { final Attributes attributes = getAttributes(); @@ -395,5 +400,13 @@ public void onReady() { listener.onReady(); } } + + @Override + public void triggerEvent(Object event) { + if (call.cancelled) { + return; + } + listener.onEvent(event); + } } } diff --git a/core/src/main/java/io/grpc/internal/ServerImpl.java b/core/src/main/java/io/grpc/internal/ServerImpl.java index d9f64c2d473..767d85f443b 100644 --- a/core/src/main/java/io/grpc/internal/ServerImpl.java +++ b/core/src/main/java/io/grpc/internal/ServerImpl.java @@ -781,6 +781,9 @@ public void closed(Status status) {} @Override public void onReady() {} + + @Override + public void triggerEvent(Object event) {} } /** @@ -960,6 +963,34 @@ public void runInContext() { callExecutor.execute(new OnReady()); } } + + @Override + public void triggerEvent(final Object event) { + try (TaskCloseable ignore = PerfMark.traceTask("ServerStreamListener.triggerEvent")) { + PerfMark.attachTag(tag); + final Link link = PerfMark.linkOut(); + + final class TriggerEvent extends ContextRunnable { + TriggerEvent() { + super(context); + } + + @Override + public void runInContext() { + try (TaskCloseable ignore = PerfMark.traceTask("ServerCallListener(app).onEvent")) { + PerfMark.attachTag(tag); + PerfMark.linkIn(link); + getListener().triggerEvent(event); + } catch (Throwable t) { + internalClose(t); + throw t; + } + } + } + + callExecutor.execute(new TriggerEvent()); + } + } } @VisibleForTesting diff --git a/core/src/main/java/io/grpc/internal/ServerStream.java b/core/src/main/java/io/grpc/internal/ServerStream.java index aa5ba10329c..4c88e94e4d7 100644 --- a/core/src/main/java/io/grpc/internal/ServerStream.java +++ b/core/src/main/java/io/grpc/internal/ServerStream.java @@ -87,6 +87,12 @@ public interface ServerStream extends Stream { */ void setListener(ServerStreamListener serverStreamListener); + /** + * Triggers a custom event. Implementations must ensure this is propagated to the + * listener on the transport thread. + */ + void triggerEvent(Object event); + /** * The context for recording stats and traces for this stream. */ diff --git a/core/src/main/java/io/grpc/internal/ServerStreamListener.java b/core/src/main/java/io/grpc/internal/ServerStreamListener.java index e55217ab422..74de0f2079e 100644 --- a/core/src/main/java/io/grpc/internal/ServerStreamListener.java +++ b/core/src/main/java/io/grpc/internal/ServerStreamListener.java @@ -42,4 +42,9 @@ public interface ServerStreamListener extends StreamListener { * @param status details about the remote closure */ void closed(Status status); + + /** + * Propagates a custom event to the listener. Must be called on the transport thread. + */ + void triggerEvent(Object event); } diff --git a/core/src/test/java/io/grpc/internal/AbstractServerStreamTest.java b/core/src/test/java/io/grpc/internal/AbstractServerStreamTest.java index 137ba19bfea..5defd17fdd0 100644 --- a/core/src/test/java/io/grpc/internal/AbstractServerStreamTest.java +++ b/core/src/test/java/io/grpc/internal/AbstractServerStreamTest.java @@ -361,6 +361,31 @@ public void close_sendTrailersClearsReservedFields() { assertEquals("bad", metadataCaptor.getValue().get(InternalStatus.MESSAGE_KEY)); } + @Test + public void triggerEvent_propagatesToListener() { + ServerStreamListener listener = mock(ServerStreamListener.class); + stream.transportState().setListener(listener); + + Object event = new Object(); + stream.triggerEvent(event); + + verify(listener).triggerEvent(event); + } + + @Test + public void triggerEvent_ignoredAfterClose() { + ServerStreamListener listener = mock(ServerStreamListener.class); + stream.transportState().setListener(listener); + + stream.close(Status.OK, new Metadata()); + stream.transportState().complete(); + + Object event = new Object(); + stream.triggerEvent(event); + + verify(listener, never()).triggerEvent(any()); + } + @Test public void changeOnReadyThreshold() { stream.setListener(new ServerStreamListenerBase()); @@ -391,6 +416,9 @@ public void halfClosed() {} @Override public void closed(Status status) {} + + @Override + public void triggerEvent(Object event) {} } private static class AbstractServerStreamBase extends AbstractServerStream { diff --git a/core/src/test/java/io/grpc/internal/ServerCallImplTest.java b/core/src/test/java/io/grpc/internal/ServerCallImplTest.java index 7394c83eab2..4a2de9f3936 100644 --- a/core/src/test/java/io/grpc/internal/ServerCallImplTest.java +++ b/core/src/test/java/io/grpc/internal/ServerCallImplTest.java @@ -493,6 +493,32 @@ public void streamListener_unexpectedRuntimeException() { assertThat(e).hasMessageThat().isEqualTo("unexpected exception"); } + @Test + public void triggerEvent_propagatesToStream() { + Object event = new Object(); + call.triggerEvent(event); + verify(stream).triggerEvent(event); + } + + @Test + public void streamListener_triggerEvent() { + ServerStreamListenerImpl streamListener = + new ServerCallImpl.ServerStreamListenerImpl<>(call, callListener, context); + Object event = new Object(); + streamListener.triggerEvent(event); + verify(callListener).onEvent(event); + } + + @Test + public void streamListener_triggerEvent_cancelled() { + ServerStreamListenerImpl streamListener = + new ServerCallImpl.ServerStreamListenerImpl<>(call, callListener, context); + Object event = new Object(); + streamListener.closed(Status.CANCELLED); + streamListener.triggerEvent(event); + verify(callListener, never()).onEvent(event); + } + private static class LongMarshaller implements Marshaller { @Override public InputStream stream(Long value) { diff --git a/core/src/test/java/io/grpc/internal/ServerImplTest.java b/core/src/test/java/io/grpc/internal/ServerImplTest.java index 91969dd6910..da2e9646042 100644 --- a/core/src/test/java/io/grpc/internal/ServerImplTest.java +++ b/core/src/test/java/io/grpc/internal/ServerImplTest.java @@ -1689,6 +1689,52 @@ public void onReady_runtimeExceptionCancelsCall() { } } + @Test + public void triggerEvent_delegatesToListener() { + JumpToApplicationThreadServerStreamListener listener + = new JumpToApplicationThreadServerStreamListener( + executor.getScheduledExecutorService(), + executor.getScheduledExecutorService(), + stream, + Context.ROOT.withCancellation(), + PerfMark.createTag()); + ServerStreamListener mockListener = mock(ServerStreamListener.class); + listener.setListener(mockListener); + + Object event = new Object(); + listener.triggerEvent(event); + + verify(mockListener, never()).triggerEvent(any()); + + executor.runDueTasks(); + verify(mockListener).triggerEvent(event); + } + + @Test + public void triggerEvent_errorCancelsCall() { + JumpToApplicationThreadServerStreamListener listener + = new JumpToApplicationThreadServerStreamListener( + executor.getScheduledExecutorService(), + executor.getScheduledExecutorService(), + stream, + Context.ROOT.withCancellation(), + PerfMark.createTag()); + ServerStreamListener mockListener = mock(ServerStreamListener.class); + listener.setListener(mockListener); + + TestError expectedT = new TestError(); + doThrow(expectedT).when(mockListener).triggerEvent(any()); + + listener.triggerEvent(new Object()); + try { + executor.runDueTasks(); + fail("Expected exception"); + } catch (TestError t) { + assertSame(expectedT, t); + ensureServerStateNotLeaked(); + } + } + @Test public void binaryLogInstalled() throws Exception { final SettableFuture intercepted = SettableFuture.create(); diff --git a/core/src/testFixtures/java/io/grpc/internal/AbstractTransportTest.java b/core/src/testFixtures/java/io/grpc/internal/AbstractTransportTest.java index 5d07de32df9..3898eb8be30 100644 --- a/core/src/testFixtures/java/io/grpc/internal/AbstractTransportTest.java +++ b/core/src/testFixtures/java/io/grpc/internal/AbstractTransportTest.java @@ -2088,6 +2088,67 @@ public void clientChecksInboundMetadataSize_trailer() throws Exception { assertNull(metadata.get(tellTaleKey)); } + @Test + public void serverStream_triggerEvent() throws Exception { + server.start(serverListener); + client = newClientTransport(server); + startTransport(client, mockClientTransportListener); + MockServerTransportListener serverTransportListener + = serverListener.takeListenerOrFail(TIMEOUT_MS, TimeUnit.MILLISECONDS); + serverTransport = serverTransportListener.transport; + + ClientStream clientStream = client.newStream( + methodDescriptor, new Metadata(), callOptions, noopTracers); + ClientStreamListenerBase clientStreamListener = new ClientStreamListenerBase(); + clientStream.start(clientStreamListener); + + StreamCreation serverStreamCreation + = serverTransportListener.takeStreamOrFail(TIMEOUT_MS, TimeUnit.MILLISECONDS); + ServerStream serverStream = serverStreamCreation.stream; + ServerStreamListenerBase serverStreamListener = serverStreamCreation.listener; + + Object event = new Object(); + serverStream.triggerEvent(event); + + Object receivedEvent = serverStreamListener.eventQueue.poll(TIMEOUT_MS, TimeUnit.MILLISECONDS); + assertEquals(event, receivedEvent); + + // Cleanup + clientStream.cancel(Status.CANCELLED); + } + + @Test + public void serverStream_triggerEvent_afterClose() throws Exception { + server.start(serverListener); + client = newClientTransport(server); + startTransport(client, mockClientTransportListener); + MockServerTransportListener serverTransportListener + = serverListener.takeListenerOrFail(TIMEOUT_MS, TimeUnit.MILLISECONDS); + serverTransport = serverTransportListener.transport; + + ClientStream clientStream = client.newStream( + methodDescriptor, new Metadata(), callOptions, noopTracers); + ClientStreamListenerBase clientStreamListener = new ClientStreamListenerBase(); + clientStream.start(clientStreamListener); + + StreamCreation serverStreamCreation + = serverTransportListener.takeStreamOrFail(TIMEOUT_MS, TimeUnit.MILLISECONDS); + ServerStream serverStream = serverStreamCreation.stream; + ServerStreamListenerBase serverStreamListener = serverStreamCreation.listener; + + // Close the stream from client side + clientStream.cancel(Status.CANCELLED); + + serverStreamListener.awaitClose(TIMEOUT_MS, TimeUnit.MILLISECONDS); + + Object event = new Object(); + serverStream.triggerEvent(event); + + // Verify listener did NOT receive the event + Object receivedEvent = serverStreamListener.eventQueue.poll(100, TimeUnit.MILLISECONDS); + assertNull(receivedEvent); + } + /** * Helper that simply does an RPC. It can be used similar to a sleep for negative testing: to give * time for actions _not_ to happen. Since it is based on doing an actual RPC with actual diff --git a/core/src/testFixtures/java/io/grpc/internal/ServerStreamListenerBase.java b/core/src/testFixtures/java/io/grpc/internal/ServerStreamListenerBase.java index aaa70600542..e4ac01912e4 100644 --- a/core/src/testFixtures/java/io/grpc/internal/ServerStreamListenerBase.java +++ b/core/src/testFixtures/java/io/grpc/internal/ServerStreamListenerBase.java @@ -89,6 +89,8 @@ public void halfClosed() { halfClosedLatch.countDown(); } + public final BlockingQueue eventQueue = new LinkedBlockingQueue<>(); + @Override public void closed(Status status) { if (this.status.isDone()) { @@ -96,4 +98,12 @@ public void closed(Status status) { } this.status.set(status); } + + @Override + public void triggerEvent(Object event) { + if (this.status.isDone()) { + fail("triggerEvent invoked after closed"); + } + eventQueue.add(event); + } } diff --git a/inprocess/src/main/java/io/grpc/inprocess/InProcessTransport.java b/inprocess/src/main/java/io/grpc/inprocess/InProcessTransport.java index a92f10fd5c5..57820b396ad 100644 --- a/inprocess/src/main/java/io/grpc/inprocess/InProcessTransport.java +++ b/inprocess/src/main/java/io/grpc/inprocess/InProcessTransport.java @@ -429,6 +429,11 @@ public void setListener(ServerStreamListener serverStreamListener) { clientStream.setListener(serverStreamListener); } + @Override + public void triggerEvent(Object event) { + clientStream.triggerServerEvent(event); + } + @Override public void request(int numMessages) { boolean onReady = clientStream.serverRequested(numMessages); @@ -732,6 +737,20 @@ private synchronized void setListener(ServerStreamListener listener) { this.serverStreamListener = listener; } + void triggerServerEvent(final Object event) { + synchronized (this) { + if (!closed) { + syncContext.executeLater(new Runnable() { + @Override + public void run() { + serverStreamListener.triggerEvent(event); + } + }); + } + } + syncContext.drain(); + } + @Override public void request(int numMessages) { boolean onReady = serverStream.clientRequested(numMessages); diff --git a/netty/src/test/java/io/grpc/netty/NettyClientTransportTest.java b/netty/src/test/java/io/grpc/netty/NettyClientTransportTest.java index ef8d2e5efda..b22be2460f1 100644 --- a/netty/src/test/java/io/grpc/netty/NettyClientTransportTest.java +++ b/netty/src/test/java/io/grpc/netty/NettyClientTransportTest.java @@ -1343,6 +1343,10 @@ public void halfClosed() { @Override public void closed(Status status) { } + + @Override + public void triggerEvent(Object event) { + } } private final class EchoServerListener implements ServerListener { diff --git a/okhttp/src/test/java/io/grpc/okhttp/OkHttpServerTransportTest.java b/okhttp/src/test/java/io/grpc/okhttp/OkHttpServerTransportTest.java index 1456f421156..ae7c53c7151 100644 --- a/okhttp/src/test/java/io/grpc/okhttp/OkHttpServerTransportTest.java +++ b/okhttp/src/test/java/io/grpc/okhttp/OkHttpServerTransportTest.java @@ -1522,6 +1522,10 @@ public void closed(Status status) { public void onReady() { } + @Override + public void triggerEvent(Object event) { + } + static String getContent(InputStream message) throws IOException { try { return new String(ByteStreams.toByteArray(message), UTF_8); diff --git a/xds/src/main/java/io/grpc/xds/ExternalProcessorFilter.java b/xds/src/main/java/io/grpc/xds/ExternalProcessorFilter.java index db7007a291b..fefe56e6ef2 100644 --- a/xds/src/main/java/io/grpc/xds/ExternalProcessorFilter.java +++ b/xds/src/main/java/io/grpc/xds/ExternalProcessorFilter.java @@ -31,6 +31,7 @@ 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.ClientInterceptor; +import io.grpc.ServerInterceptor; import io.grpc.internal.GrpcUtil; import io.grpc.xds.internal.HeaderForwardingRulesConfig; import io.grpc.xds.internal.grpcservice.CachedChannelManager; @@ -81,6 +82,11 @@ public boolean isClientFilter() { return GrpcUtil.getFlag("GRPC_EXPERIMENTAL_XDS_EXT_PROC_ON_CLIENT", false); } + @Override + public boolean isServerFilter() { + return GrpcUtil.getFlag("GRPC_EXPERIMENTAL_XDS_EXT_PROC_ON_SERVER", false); + } + @Override public ExternalProcessorFilter newInstance(FilterContext context) { return new ExternalProcessorFilter(context); @@ -134,6 +140,20 @@ public ClientInterceptor buildClientInterceptor(FilterConfig filterConfig, extProcFilterConfig, cachedChannelManager, scheduler, context); } + @Nullable + @Override + public ServerInterceptor buildServerInterceptor(FilterConfig filterConfig, + @Nullable FilterConfig overrideConfig) { + ExternalProcessorFilterConfig extProcFilterConfig = + (ExternalProcessorFilterConfig) filterConfig; + if (overrideConfig != null) { + extProcFilterConfig = mergeConfigs(extProcFilterConfig, + (ExternalProcessorFilterOverrideConfig) overrideConfig); + } + return new ExternalProcessorServerInterceptor( + extProcFilterConfig, cachedChannelManager, context); + } + private static ExternalProcessorFilterConfig mergeConfigs( ExternalProcessorFilterConfig extProcFilterConfig, ExternalProcessorFilterOverrideConfig extProcFilterConfigOverride) { diff --git a/xds/src/main/java/io/grpc/xds/ExternalProcessorServerInterceptor.java b/xds/src/main/java/io/grpc/xds/ExternalProcessorServerInterceptor.java new file mode 100644 index 00000000000..9d0046291c4 --- /dev/null +++ b/xds/src/main/java/io/grpc/xds/ExternalProcessorServerInterceptor.java @@ -0,0 +1,1915 @@ +/* + * Copyright 2024 The gRPC Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package io.grpc.xds; + +import static com.google.common.base.Preconditions.checkNotNull; +import static io.grpc.xds.internal.extproc.ExternalProcessorUtil.applyHeaderMutations; +import static io.grpc.xds.internal.extproc.ExternalProcessorUtil.collectAttributes; +import static io.grpc.xds.internal.extproc.ExternalProcessorUtil.markDataPlaneCallClosed; +import static io.grpc.xds.internal.extproc.ExternalProcessorUtil.markExtProcStreamCompleted; +import static io.grpc.xds.internal.extproc.ExternalProcessorUtil.markExtProcStreamFailed; +import static io.grpc.xds.internal.extproc.ExternalProcessorUtil.outboundStreamToByteString; +import static io.grpc.xds.internal.extproc.ExternalProcessorUtil.toHeaderMap; + +import com.google.common.annotations.VisibleForTesting; +import com.google.common.collect.ImmutableList; +import com.google.common.collect.ImmutableMap; +import com.google.common.util.concurrent.MoreExecutors; +import com.google.protobuf.ByteString; +import com.google.protobuf.Struct; +import io.envoyproxy.envoy.extensions.filters.http.ext_proc.v3.ProcessingMode; +import io.envoyproxy.envoy.service.ext_proc.v3.BodyMutation; +import io.envoyproxy.envoy.service.ext_proc.v3.BodyResponse; +import io.envoyproxy.envoy.service.ext_proc.v3.CommonResponse; +import io.envoyproxy.envoy.service.ext_proc.v3.ExternalProcessorGrpc; +import io.envoyproxy.envoy.service.ext_proc.v3.HttpBody; +import io.envoyproxy.envoy.service.ext_proc.v3.HttpHeaders; +import io.envoyproxy.envoy.service.ext_proc.v3.HttpTrailers; +import io.envoyproxy.envoy.service.ext_proc.v3.ImmediateResponse; +import io.envoyproxy.envoy.service.ext_proc.v3.ProcessingRequest; +import io.envoyproxy.envoy.service.ext_proc.v3.ProcessingResponse; +import io.envoyproxy.envoy.service.ext_proc.v3.ProtocolConfiguration; +import io.envoyproxy.envoy.service.ext_proc.v3.StreamedBodyResponse; +import io.grpc.Context; +import io.grpc.DoubleHistogramMetricInstrument; +import io.grpc.ForwardingServerCall.SimpleForwardingServerCall; +import io.grpc.ManagedChannel; +import io.grpc.Metadata; +import io.grpc.MethodDescriptor; +import io.grpc.MetricInstrumentRegistry; +import io.grpc.MetricRecorder; +import io.grpc.ServerCall; +import io.grpc.ServerCallHandler; +import io.grpc.ServerInterceptor; +import io.grpc.Status; +import io.grpc.StatusRuntimeException; +import io.grpc.SynchronizationContext; +import io.grpc.internal.GrpcUtil; +import io.grpc.internal.SharedResourceHolder; +import io.grpc.stub.ClientCallStreamObserver; +import io.grpc.stub.ClientResponseObserver; +import io.grpc.stub.MetadataUtils; +import io.grpc.xds.ExternalProcessorFilter.ExternalProcessorFilterConfig; +import io.grpc.xds.Filter.FilterContext; +import io.grpc.xds.internal.extproc.DataPlaneCallState; +import io.grpc.xds.internal.extproc.EventType; +import io.grpc.xds.internal.extproc.ExtProcStreamState; +import io.grpc.xds.internal.extproc.KnownLengthInputStream; +import io.grpc.xds.internal.grpcservice.CachedChannelManager; +import io.grpc.xds.internal.grpcservice.HeaderValue; +import io.grpc.xds.internal.headermutations.HeaderMutationDisallowedException; +import io.grpc.xds.internal.headermutations.HeaderMutationFilter; +import io.grpc.xds.internal.headermutations.HeaderMutationRulesConfig; +import io.grpc.xds.internal.headermutations.HeaderMutator; +import java.io.IOException; +import java.io.InputStream; +import java.util.List; +import java.util.ArrayList; +import java.util.Optional; +import java.util.Queue; +import java.util.concurrent.ConcurrentLinkedQueue; +import java.util.concurrent.ScheduledExecutorService; +import java.util.concurrent.ScheduledFuture; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; +import java.util.logging.Level; +import java.util.logging.Logger; +import javax.annotation.Nullable; +import javax.annotation.concurrent.GuardedBy; + +/** + * Server-side interceptor for external processing filter. + */ +final class ExternalProcessorServerInterceptor implements ServerInterceptor { + private static final Logger logger = Logger.getLogger( + ExternalProcessorServerInterceptor.class.getName()); + + @VisibleForTesting + static DoubleHistogramMetricInstrument clientHeadersDuration; + @VisibleForTesting + static DoubleHistogramMetricInstrument clientHalfCloseDuration; + @VisibleForTesting + static DoubleHistogramMetricInstrument serverHeadersDuration; + @VisibleForTesting + static DoubleHistogramMetricInstrument serverTrailersDuration; + + // Copied from io.grpc.opentelemetry.internal.OpenTelemetryConstants.LATENCY_BUCKETS + private static final List LATENCY_BUCKETS = ImmutableList.of( + 0d, 0.00001d, 0.00005d, 0.0001d, 0.0003d, 0.0006d, 0.0008d, 0.001d, 0.002d, + 0.003d, 0.004d, 0.005d, 0.006d, 0.008d, 0.01d, 0.013d, 0.016d, 0.02d, + 0.025d, 0.03d, 0.04d, 0.05d, 0.065d, 0.08d, 0.1d, 0.13d, 0.16d, + 0.2d, 0.25d, 0.3d, 0.4d, 0.5d, 0.65d, 0.8d, 1d, 2d, + 5d, 10d, 20d, 50d, 100d); + + static { + initMetricInstruments(); + } + + public static synchronized void initMetricInstruments() { + if (GrpcUtil.getFlag("GRPC_EXPERIMENTAL_XDS_EXT_PROC_ON_SERVER", false)) { + if (clientHeadersDuration == null) { + MetricInstrumentRegistry registry = MetricInstrumentRegistry.getDefaultRegistry(); + + clientHeadersDuration = registry.registerDoubleHistogram( + "grpc.server_ext_proc.client_headers_duration", + "Time between when the ext_proc filter sees the client's headers and when " + + "it allows those headers to continue on to the next filter", + "s", + LATENCY_BUCKETS, + ImmutableList.of(), + ImmutableList.of(), + true); + + clientHalfCloseDuration = registry.registerDoubleHistogram( + "grpc.server_ext_proc.client_half_close_duration", + "Time between when the ext_proc filter sees the client's half-close and when " + + "it allows that half-close to continue on to the next filter", + "s", + LATENCY_BUCKETS, + ImmutableList.of(), + ImmutableList.of(), + true); + + serverHeadersDuration = registry.registerDoubleHistogram( + "grpc.server_ext_proc.server_headers_duration", + "Time between when the ext_proc filter sees the server's headers and when " + + "it allows those headers to continue on to the next filter", + "s", + LATENCY_BUCKETS, + ImmutableList.of(), + ImmutableList.of(), + true); + + serverTrailersDuration = registry.registerDoubleHistogram( + "grpc.server_ext_proc.server_trailers_duration", + "Time between when the ext_proc filter sees the server's trailers and when " + + "it allows those trailers to continue on to the next filter", + "s", + LATENCY_BUCKETS, + ImmutableList.of(), + ImmutableList.of(), + true); + } + } + } + + private final ExternalProcessorFilterConfig filterConfig; + private final MetricRecorder metricsRecorder; + private final ManagedChannel extProcChannel; + + ExternalProcessorServerInterceptor( + ExternalProcessorFilterConfig filterConfig, + CachedChannelManager cachedChannelManager, + FilterContext context) { + this.filterConfig = checkNotNull(filterConfig, "filterConfig"); + checkNotNull(cachedChannelManager, "cachedChannelManager"); + this.metricsRecorder = checkNotNull(context.metricsRecorder(), "metricsRecorder"); + this.extProcChannel = cachedChannelManager.getChannel(filterConfig.getGrpcServiceConfig()); + } + + ExternalProcessorFilterConfig getFilterConfig() { + return filterConfig; + } + + @Override + @SuppressWarnings("unchecked") + public ServerCall.Listener interceptCall( + ServerCall call, Metadata headers, ServerCallHandler next) { + ServerCall rawCall = (ServerCall) call; + ServerCallHandler rawNext = + (ServerCallHandler) next; + + ExternalProcessorGrpc.ExternalProcessorStub extProcStub = ExternalProcessorGrpc.newStub( + extProcChannel) + .withExecutor(MoreExecutors.directExecutor()); + + if (filterConfig.getGrpcServiceConfig().timeout().isPresent()) { + long timeoutNanos = filterConfig.getGrpcServiceConfig().timeout().get().toNanos(); + if (timeoutNanos > 0) { + extProcStub = extProcStub.withDeadlineAfter(timeoutNanos, TimeUnit.NANOSECONDS); + } + } + if (filterConfig.getGrpcServiceConfig().initialMetadata() != null + && !filterConfig.getGrpcServiceConfig().initialMetadata().isEmpty()) { + Metadata extraHeaders = new Metadata(); + for (HeaderValue headerValue : filterConfig.getGrpcServiceConfig().initialMetadata()) { + String key = headerValue.key(); + if (key.endsWith(Metadata.BINARY_HEADER_SUFFIX)) { + if (headerValue.rawValue().isPresent()) { + Metadata.Key metadataKey = + Metadata.Key.of(key, Metadata.BINARY_BYTE_MARSHALLER); + extraHeaders.put(metadataKey, headerValue.rawValue().get().toByteArray()); + } + } else { + if (headerValue.value().isPresent()) { + Metadata.Key metadataKey = + Metadata.Key.of(key, Metadata.ASCII_STRING_MARSHALLER); + extraHeaders.put(metadataKey, headerValue.value().get()); + } + } + } + extProcStub = extProcStub.withInterceptors( + MetadataUtils.newAttachHeadersInterceptor(extraHeaders)); + } + + Context callContext = Context.current(); + + DataPlaneServerCall dataPlaneServerCall = new DataPlaneServerCall( + rawCall, extProcStub, filterConfig, filterConfig.getMutationRulesConfig(), + SharedResourceHolder.get(GrpcUtil.TIMER_SERVICE), call.getMethodDescriptor(), + metricsRecorder, call.getAuthority(), rawNext, headers, callContext); + + dataPlaneServerCall.start(); + + return (ServerCall.Listener) dataPlaneServerCall.getListener(); + } + + + static class DataPlaneServerCall + extends SimpleForwardingServerCall { + + private final ServerCall rawCall; + private final ExternalProcessorGrpc.ExternalProcessorStub extProcStub; + private final ExternalProcessorFilterConfig config; + private final ScheduledExecutorService scheduler; + final Object streamLock = new Object(); + private final Queue expectedResponses = new ConcurrentLinkedQueue<>(); + private volatile ClientCallStreamObserver extProcClientCallRequestObserver; + private final Queue pendingDrainingMessages = new ConcurrentLinkedQueue<>(); + private final Queue savedOutgoingMessages = new ConcurrentLinkedQueue<>(); + private volatile DataPlaneServerListener wrappedListener; + private final HeaderMutationFilter mutationFilter; + private final HeaderMutator mutator = HeaderMutator.create(); + private final AtomicInteger pendingRequests = new AtomicInteger(0); + private final ProcessingMode currentProcessingMode; + private final MethodDescriptor method; + private final MetricRecorder metricsRecorder; + private final String authority; + private final ServerCallHandler rawNext; + private final Context callContext; + private volatile Metadata requestHeaders; + + @GuardedBy("streamLock") + private volatile Metadata savedResponseHeaders; + @GuardedBy("streamLock") + private volatile Status savedStatus; + @GuardedBy("streamLock") + private volatile Metadata savedTrailers; + + @GuardedBy("streamLock") + private boolean protocolConfigSent = false; + @GuardedBy("streamLock") + private ImmutableMap collectedAttributes; + @GuardedBy("streamLock") + private boolean requestAttributesSent = false; + + // Default initial window size + private static final long DEFAULT_INITIAL_WINDOW_SIZE = 65536; + + // Outbound (sending) windows + @GuardedBy("streamLock") + private long downstreamToSidestreamWindow = DEFAULT_INITIAL_WINDOW_SIZE; + @GuardedBy("streamLock") + private long upstreamToSidestreamWindow = DEFAULT_INITIAL_WINDOW_SIZE; + + // Inbound (receiving) windows + @GuardedBy("streamLock") + private long sidestreamToUpstreamWindow = DEFAULT_INITIAL_WINDOW_SIZE; + @GuardedBy("streamLock") + private long sidestreamToDownstreamWindow = DEFAULT_INITIAL_WINDOW_SIZE; + + // Path 1: Pending/buffered request body messages from client to send to ext_proc + @GuardedBy("streamLock") + final Queue pendingRequestBodyMessages = new ConcurrentLinkedQueue<>(); + + // Path 2: Buffered mutated request bodies from ext_proc to deliver to App + @GuardedBy("streamLock") + final Queue pendingMutatedRequestBodies = new ConcurrentLinkedQueue<>(); + @GuardedBy("streamLock") + private final AtomicInteger pendingAppRequests = new AtomicInteger(0); + @GuardedBy("streamLock") + private final AtomicInteger pendingTransportRequests = new AtomicInteger(0); + + // Path 3: Pending/buffered response body messages from App to send to ext_proc + @GuardedBy("streamLock") + private final Queue pendingResponseBodyMessages = new ConcurrentLinkedQueue<>(); + @GuardedBy("streamLock") + private final AtomicBoolean pendingClose = new AtomicBoolean(false); + + // Path 4: Buffered mutated response bodies from ext_proc to send to client transport + @GuardedBy("streamLock") + private final Queue pendingDownstreamBodyMessages = new ConcurrentLinkedQueue<>(); + + // Threshold to trigger standalone client window updates + private static final long WINDOW_UPDATE_THRESHOLD = DEFAULT_INITIAL_WINDOW_SIZE / 2; + + // Accumulated client window updates to send to ext_proc + @GuardedBy("streamLock") + private long accumulatedWindowUpdateSidestreamToUpstream = 0; + @GuardedBy("streamLock") + private long accumulatedWindowUpdateSidestreamToDownstream = 0; + + // Flag to track if FlowControlInit was sent in the initial message + @GuardedBy("streamLock") + private boolean flowControlInitSent = false; + + private long clientHeadersStartNanos; + private long clientHalfCloseStartNanos; + private long serverHeadersStartNanos; + private long serverTrailersStartNanos; + + final AtomicReference dataPlaneCallState = + new AtomicReference<>(DataPlaneCallState.IDLE); + final AtomicReference extProcStreamState = + new AtomicReference<>(ExtProcStreamState.ACTIVE); + final AtomicBoolean passThroughMode = new AtomicBoolean(false); + final AtomicBoolean halfClosed = new AtomicBoolean(false); + final AtomicBoolean requestSideClosed = new AtomicBoolean(false); + final AtomicBoolean dataPlaneCallClosed = new AtomicBoolean(false); + final AtomicBoolean bodyMessageSentToExtProc = new AtomicBoolean(false); + + final AtomicBoolean isProcessingTrailers = new AtomicBoolean(false); + final AtomicBoolean responseHeadersSent = new AtomicBoolean(false); + final AtomicBoolean trailersOnly = new AtomicBoolean(false); + final AtomicBoolean terminationTriggered = new AtomicBoolean(false); + + protected DataPlaneServerCall( + ServerCall rawCall, + ExternalProcessorGrpc.ExternalProcessorStub extProcStub, + ExternalProcessorFilterConfig config, + Optional mutationRulesConfig, + ScheduledExecutorService scheduler, + MethodDescriptor method, + MetricRecorder metricsRecorder, + String authority, + ServerCallHandler rawNext, + Metadata requestHeaders, + Context callContext) { + super(rawCall); + this.rawCall = rawCall; + this.extProcStub = extProcStub.withExecutor(MoreExecutors.directExecutor()); + this.config = config; + this.currentProcessingMode = config.getExternalProcessor().getProcessingMode(); + this.mutationFilter = new HeaderMutationFilter(mutationRulesConfig); + this.scheduler = scheduler; + this.method = method; + this.metricsRecorder = checkNotNull(metricsRecorder, "metricsRecorder"); + this.authority = authority; + this.rawNext = rawNext; + this.requestHeaders = requestHeaders; + this.callContext = callContext; + this.wrappedListener = new DataPlaneServerListener(this); + } + + DataPlaneServerListener getListener() { + return wrappedListener; + } + + boolean isExtProcStreamCompleted() { + return extProcStreamState.get().isCompleted(); + } + + boolean isExtProcStreamFailed() { + return extProcStreamState.get().isFailed(); + } + + boolean isExtProcStreamDraining() { + return extProcStreamState.get().isDraining(); + } + + private boolean isSimpleServerCall(Class clazz) { + if (clazz == null) { + return false; + } + if (clazz.getName().contains("SimpleServerCall")) { + return true; + } + return isSimpleServerCall(clazz.getSuperclass()); + } + + @Override + public void triggerEvent(Object event) { + if (isSimpleServerCall(rawCall.getClass())) { + wrappedListener.onEvent(event); + } else { + super.triggerEvent(event); + } + } + + private void activateCall() { + if ((extProcStreamState.get() == ExtProcStreamState.FAILED + && !config.getFailureModeAllow() + && !config.getObservabilityMode()) + || !dataPlaneCallState.compareAndSet( + DataPlaneCallState.IDLE, DataPlaneCallState.ACTIVE)) { + return; + } + if (clientHeadersStartNanos > 0) { + long durationNanos = System.nanoTime() - clientHeadersStartNanos; + recordDuration(clientHeadersDuration, durationNanos); + clientHeadersStartNanos = 0; + } + Context previous = callContext.attach(); + ServerCall.Listener appListener; + try { + appListener = rawNext.startCall(this, requestHeaders); + } finally { + callContext.detach(previous); + } + wrappedListener.setDelegate(appListener); + drainPendingRequests(); + wrappedListener.onReadyNotify(); + if (wrappedListener.halfCloseDeferred) { + wrappedListener.handleDeferredHalfClose(); + } + } + + private void recordDuration(DoubleHistogramMetricInstrument instrument, long durationNanos) { + if (instrument != null) { + double durationSecs = (double) durationNanos / 1_000_000_000.0; + metricsRecorder.recordDoubleHistogram( + instrument, + durationSecs, + ImmutableList.of("server"), + ImmutableList.of("server")); + } + } + + private boolean validateCompressionSupport(BodyResponse bodyResponse) { + if (bodyResponse.hasResponse() && bodyResponse.getResponse().hasBodyMutation()) { + BodyMutation mutation = bodyResponse.getResponse().getBodyMutation(); + if (mutation.hasStreamedResponse() + && mutation.getStreamedResponse().getGrpcMessageCompressed()) { + StatusRuntimeException ex = Status.UNAVAILABLE + .withDescription("gRPC message compression not supported in ext_proc") + .asRuntimeException(); + synchronized (streamLock) { + if (!isExtProcStreamCompleted() && extProcClientCallRequestObserver != null) { + extProcClientCallRequestObserver.onError(ex); + } + } + activateCall(); + markExtProcStreamFailed(extProcStreamState); + rawCall.close( + Status.UNAVAILABLE.withDescription( + "gRPC message compression not supported in ext_proc"), + new Metadata()); + closeExtProcStream(); + return false; + } + } + return true; + } + + + + void start() { + clientHeadersStartNanos = System.nanoTime(); + synchronized (streamLock) { + this.collectedAttributes = collectAttributes( + config.getRequestAttributes(), method, authority, requestHeaders); + } + + extProcStub.process(new ClientResponseObserver() { + @Override + public void beforeStart(ClientCallStreamObserver requestStream) { + synchronized (streamLock) { + extProcClientCallRequestObserver = requestStream; + } + requestStream.setOnReadyHandler(() -> DataPlaneServerCall.this.triggerEvent(new ExtProcStreamReadyEvent())); + } + + @Override + public void onNext(ProcessingResponse response) { + DataPlaneServerCall.this.triggerEvent(new ExtProcResponseEvent(response)); + } + + @Override + public void onError(Throwable t) { + DataPlaneServerCall.this.triggerEvent(new ExtProcErrorEvent(t)); + } + + @Override + public void onCompleted() { + DataPlaneServerCall.this.triggerEvent(new ExtProcCompletedEvent()); + } + }); + + boolean sendRequestHeaders = + currentProcessingMode.getRequestHeaderMode() == ProcessingMode.HeaderSendMode.SEND + || currentProcessingMode.getRequestHeaderMode() + == ProcessingMode.HeaderSendMode.DEFAULT; + + if (sendRequestHeaders) { + sendToExtProc(ProcessingRequest.newBuilder() + .setRequestHeaders(HttpHeaders.newBuilder() + .setHeaders(toHeaderMap(requestHeaders, config.getForwardRulesConfig())) + .setEndOfStream(false) + .build()) + .build()); + } + + if (config.getObservabilityMode() || !sendRequestHeaders) { + activateCall(); + } + } + + private void sendToExtProc(ProcessingRequest request) { + synchronized (streamLock) { + if (isExtProcStreamCompleted()) { + return; + } + + ProcessingRequest requestToSend = request; + if (!protocolConfigSent) { + requestToSend = ProcessingRequest.newBuilder(requestToSend) + .setProtocolConfig(ProtocolConfiguration.newBuilder() + .setRequestBodyMode(currentProcessingMode.getRequestBodyMode()) + .setResponseBodyMode(currentProcessingMode.getResponseBodyMode()) + .build()) + .build(); + protocolConfigSent = true; + } + + boolean isClientServerMessage = + requestToSend.hasRequestHeaders() || requestToSend.hasRequestBody(); + if (isClientServerMessage + && !requestAttributesSent + && collectedAttributes != null + && !collectedAttributes.isEmpty()) { + requestToSend = ProcessingRequest.newBuilder(requestToSend) + .putAllAttributes(collectedAttributes) + .build(); + requestAttributesSent = true; + } + + if (config.getObservabilityMode()) { + requestToSend = ProcessingRequest.newBuilder(requestToSend) + .setObservabilityMode(true) + .build(); + } else if (!flowControlInitSent) { + requestToSend = ProcessingRequest.newBuilder(requestToSend) + .setFlowControlInit(ProcessingRequest.FlowControlInit.newBuilder() + .setInitialWindowDownstreamToSidestream(DEFAULT_INITIAL_WINDOW_SIZE) + .setInitialWindowSidestreamToUpstream(DEFAULT_INITIAL_WINDOW_SIZE) + .setInitialWindowUpstreamToSidestreama(DEFAULT_INITIAL_WINDOW_SIZE) + .setInitialWindowSidestreamToDownstream(DEFAULT_INITIAL_WINDOW_SIZE) + .build()) + .build(); + flowControlInitSent = true; + } + + if (requestToSend.hasRequestHeaders()) { + expectedResponses.add(EventType.REQUEST_HEADERS); + } else if (requestToSend.hasRequestBody()) { + expectedResponses.add(EventType.REQUEST_BODY); + } else if (requestToSend.hasResponseHeaders()) { + expectedResponses.add(EventType.RESPONSE_HEADERS); + } else if (requestToSend.hasResponseBody()) { + expectedResponses.add(EventType.RESPONSE_BODY); + } else if (requestToSend.hasResponseTrailers()) { + expectedResponses.add(EventType.RESPONSE_TRAILERS); + } + + extProcClientCallRequestObserver.onNext(requestToSend); + } + } + + @GuardedBy("streamLock") + private void mergeAccumulatedWindowUpdates(ProcessingRequest.Builder requestBuilder) { + long incrementUpstream = accumulatedWindowUpdateSidestreamToUpstream; + long incrementDownstream = accumulatedWindowUpdateSidestreamToDownstream; + + if (incrementUpstream > 0 || incrementDownstream > 0) { + requestBuilder.setClientWindowUpdate( + ProcessingRequest.ClientWindowUpdate.newBuilder() + .setWindowIncrementSidestreamToUpstream(incrementUpstream) + .setWindowIncrementSidestreamToDownstream(incrementDownstream) + .build()); + accumulatedWindowUpdateSidestreamToUpstream -= incrementUpstream; + accumulatedWindowUpdateSidestreamToDownstream -= incrementDownstream; + sidestreamToUpstreamWindow += incrementUpstream; + sidestreamToDownstreamWindow += incrementDownstream; + } + } + + private void trySendAccumulatedWindowUpdates() { + synchronized (streamLock) { + if (isExtProcStreamCompleted()) { + return; + } + if (accumulatedWindowUpdateSidestreamToUpstream >= WINDOW_UPDATE_THRESHOLD + || accumulatedWindowUpdateSidestreamToDownstream >= WINDOW_UPDATE_THRESHOLD + || sidestreamToUpstreamWindow < 0 + || sidestreamToDownstreamWindow < 0) { + ProcessingRequest.Builder builder = ProcessingRequest.newBuilder(); + mergeAccumulatedWindowUpdates(builder); + if (builder.hasClientWindowUpdate()) { + sendToExtProc(builder.build()); + } + } + } + } + + void onExtProcStreamReady() { + drainPendingRequests(); + wrappedListener.onReadyNotify(); + } + + private void drainPendingRequests() { + if (config.getObservabilityMode() + || currentProcessingMode.getRequestBodyMode() != ProcessingMode.BodySendMode.GRPC + || isExtProcStreamCompleted()) { + int toRequest = pendingRequests.getAndSet(0); + if (toRequest > 0) { + super.request(toRequest); + } + return; + } + + // Normal mode flow control: pull 1 message at a time + while (true) { + boolean pull = false; + synchronized (streamLock) { + if (isSidecarReady() + && downstreamToSidestreamWindow > 0 + && pendingTransportRequests.get() > 0) { + pull = true; + pendingTransportRequests.decrementAndGet(); + } + } + if (pull) { + super.request(1); + } else { + break; + } + } + } + + private void closeExtProcStream() { + synchronized (streamLock) { + if (markExtProcStreamCompleted(extProcStreamState)) { + if (extProcClientCallRequestObserver != null) { + extProcClientCallRequestObserver.onCompleted(); + } + } + expectedResponses.clear(); + } + proceedWithClose(); + } + + private void cancelExtProcStream(Throwable t) { + if (markExtProcStreamFailed(extProcStreamState)) { + synchronized (streamLock) { + if (extProcClientCallRequestObserver != null) { + try { + extProcClientCallRequestObserver.onError(t); + } catch (Throwable ignored) { + // Ignore exceptions during cancel/onError propagation + } + extProcClientCallRequestObserver = null; + } + } + expectedResponses.clear(); + proceedWithClose(); + } + } + + private void internalOnError(Throwable t) { + if (markExtProcStreamFailed(extProcStreamState)) { + synchronized (streamLock) { + if (extProcClientCallRequestObserver != null) { + try { + extProcClientCallRequestObserver.onError(t); + } catch (Throwable ignored) { + // Ignore exceptions during cancel/onError propagation + } + extProcClientCallRequestObserver = null; + } + } + expectedResponses.clear(); + if (config.getObservabilityMode() + || (config.getFailureModeAllow() && !bodyMessageSentToExtProc.get())) { + handleFailOpen(); + } else { + proceedWithClose( + Status.INTERNAL.withDescription("External processor stream failed").withCause(t), + new Metadata()); + } + } + } + + void handleExtProcResponse(ProcessingResponse response) { + try { + if (config.getObservabilityMode()) { + return; + } + + if (response.hasServerWindowUpdate()) { + ProcessingResponse.ServerWindowUpdate update = response.getServerWindowUpdate(); + synchronized (streamLock) { + downstreamToSidestreamWindow += update.getWindowIncrementDownstreamToSidestream(); + upstreamToSidestreamWindow += update.getWindowIncrementUpstreamToSidestream(); + } + if (wrappedListener != null) { + wrappedListener.drainPendingRequestBodyMessages(); + } + drainPendingResponseBodyMessages(); + drainPendingRequests(); + } + + if (response.hasImmediateResponse()) { + if (config.getDisableImmediateResponse()) { + internalOnError(Status.UNAVAILABLE + .withDescription( + "Immediate response is disabled but received from external processor") + .asRuntimeException()); + return; + } + handleImmediateResponse(response.getImmediateResponse()); + return; + } + + EventType expected = expectedResponses.peek(); + EventType received = null; + if (response.hasRequestHeaders()) { + received = EventType.REQUEST_HEADERS; + } else if (response.hasRequestBody()) { + received = EventType.REQUEST_BODY; + } else if (response.hasResponseHeaders()) { + received = EventType.RESPONSE_HEADERS; + } else if (response.hasResponseBody()) { + received = EventType.RESPONSE_BODY; + } else if (response.hasResponseTrailers()) { + received = EventType.RESPONSE_TRAILERS; + } + + if (received != null) { + if (expected == null || expected != received) { + internalOnError(Status.UNAVAILABLE + .withDescription("Protocol error: received response out of order. Expected: " + + expected + ", Received: " + received) + .asRuntimeException()); + return; + } + expectedResponses.poll(); + } + + if (response.getRequestDrain()) { + extProcStreamState.set(ExtProcStreamState.DRAINING); + activateCall(); + halfCloseExtProcStream(); + } + + if (response.hasRequestHeaders()) { + if (response.getRequestHeaders().hasResponse()) { + if (response.getRequestHeaders().getResponse().getStatus() + == CommonResponse.ResponseStatus.CONTINUE_AND_REPLACE) { + internalOnError(Status.UNAVAILABLE + .withDescription("CONTINUE_AND_REPLACE is not supported") + .asRuntimeException()); + return; + } + applyHeaderMutations( + requestHeaders, + response.getRequestHeaders().getResponse().getHeaderMutation(), + mutationFilter, + mutator); + } + activateCall(); + } + else if (response.hasRequestBody()) { + if (validateCompressionSupport(response.getRequestBody())) { + handleRequestBodyResponse(response.getRequestBody()); + } + } + else if (response.hasResponseHeaders()) { + if (response.getResponseHeaders().hasResponse()) { + if (response.getResponseHeaders().getResponse().getStatus() + == CommonResponse.ResponseStatus.CONTINUE_AND_REPLACE) { + internalOnError(Status.UNAVAILABLE + .withDescription("CONTINUE_AND_REPLACE is not supported") + .asRuntimeException()); + return; + } + synchronized (streamLock) { + applyHeaderMutations( + trailersOnly.get() ? savedTrailers : savedResponseHeaders, + response.getResponseHeaders().getResponse().getHeaderMutation(), + mutationFilter, + mutator); + } + } + if (trailersOnly.get()) { + proceedWithClose(); + } else { + proceedWithSendHeaders(); + } + } + else if (response.hasResponseBody()) { + if (validateCompressionSupport(response.getResponseBody())) { + handleResponseBodyResponse(response.getResponseBody()); + } + } + else if (response.hasResponseTrailers()) { + if (response.getResponseTrailers().hasHeaderMutation()) { + synchronized (streamLock) { + applyHeaderMutations( + savedTrailers, + response.getResponseTrailers().getHeaderMutation(), + mutationFilter, + mutator); + } + } + proceedWithClose(); + } + + checkEndOfStream(); + } catch (Throwable t) { + internalOnError(t); + } + } + + void handleExtProcError(Throwable t) { + if (markExtProcStreamFailed(extProcStreamState)) { + synchronized (streamLock) { + extProcClientCallRequestObserver = null; + } + if (config.getObservabilityMode() + || (config.getFailureModeAllow() && !bodyMessageSentToExtProc.get())) { + handleFailOpen(); + } else { + proceedWithClose( + Status.INTERNAL.withDescription("External processor stream failed") + .withCause(t), + new Metadata()); + } + } + } + + void handleExtProcCompleted() { + ExtProcStreamState state = extProcStreamState.get(); + if (state == ExtProcStreamState.DRAINING) { + if (markExtProcStreamCompleted(extProcStreamState)) { + handleFailOpen(); + } + } else if (state == ExtProcStreamState.ACTIVE) { + internalOnError(Status.UNAVAILABLE + .withDescription("External processor stream completed without drain") + .asRuntimeException()); + } + } + + private void halfCloseExtProcStream() { + synchronized (streamLock) { + if (!isExtProcStreamCompleted() && extProcClientCallRequestObserver != null) { + extProcClientCallRequestObserver.onCompleted(); + } + } + } + + private boolean isSidecarReady() { + if (isExtProcStreamCompleted()) { + return true; + } + if (isExtProcStreamDraining()) { + return false; + } + synchronized (streamLock) { + ClientCallStreamObserver observer = extProcClientCallRequestObserver; + return observer != null && observer.isReady(); + } + } + + @Override + public boolean isReady() { + if (passThroughMode.get()) { + return super.isReady(); + } + if (isExtProcStreamCompleted()) { + return super.isReady(); + } + if (dataPlaneCallState.get() == DataPlaneCallState.IDLE && !config.getObservabilityMode()) { + return false; + } + synchronized (streamLock) { + boolean sidecarReady = isSidecarReady(); + if (config.getObservabilityMode()) { + return super.isReady() && sidecarReady; + } + return upstreamToSidestreamWindow > 0 && sidecarReady + && pendingResponseBodyMessages.isEmpty(); + } + } + + @Override + public void request(int numMessages) { + if (passThroughMode.get() || isExtProcStreamCompleted()) { + super.request(numMessages); + return; + } + if (currentProcessingMode.getRequestBodyMode() != ProcessingMode.BodySendMode.GRPC) { + synchronized (streamLock) { + if (isSidecarReady()) { + super.request(numMessages); + } else { + pendingRequests.addAndGet(numMessages); + } + } + return; + } + if (config.getObservabilityMode()) { + synchronized (streamLock) { + if (isSidecarReady()) { + super.request(numMessages); + } else { + pendingRequests.addAndGet(numMessages); + } + } + return; + } + + synchronized (streamLock) { + pendingAppRequests.addAndGet(numMessages); + } + + int satisfied = drainPendingMutatedRequestBodies(); + + synchronized (streamLock) { + int remaining = numMessages - satisfied; + if (remaining > 0) { + pendingTransportRequests.addAndGet(remaining); + } + } + + drainPendingRequests(); + } + + @Override + public void sendHeaders(Metadata headers) { + + serverHeadersStartNanos = System.nanoTime(); + responseHeadersSent.set(true); + boolean sendResponseHeaders = + currentProcessingMode.getResponseHeaderMode() == ProcessingMode.HeaderSendMode.SEND + || currentProcessingMode.getResponseHeaderMode() + == ProcessingMode.HeaderSendMode.DEFAULT; + + synchronized (streamLock) { + // NOTE: Even if sendResponseHeaders is false, we MUST obtain streamLock to call + // proceedWithSendHeaders() safely, because an active control plane thread could + // concurrently call super.sendMessage() or super.close() (e.g., due to a concurrent error). + if (passThroughMode.get() || isExtProcStreamCompleted() || !sendResponseHeaders) { + proceedWithSendHeaders(headers); + return; + } + this.savedResponseHeaders = headers; + if (isExtProcStreamDraining()) { + return; + } + } + + sendToExtProc(ProcessingRequest.newBuilder() + .setResponseHeaders(HttpHeaders.newBuilder() + .setHeaders(toHeaderMap(headers, config.getForwardRulesConfig())) + .build()) + .build()); + + if (config.getObservabilityMode()) { + synchronized (streamLock) { + proceedWithSendHeaders(); + } + } + } + + void proceedWithSendHeaders() { + synchronized (streamLock) { + if (savedResponseHeaders != null) { + proceedWithSendHeaders(savedResponseHeaders); + savedResponseHeaders = null; + InputStream msg; + while ((msg = savedOutgoingMessages.poll()) != null) { + sendMessage(msg); + } + if (savedStatus != null) { + triggerCloseHandshake(savedTrailers); + } + } + } + } + + private void proceedWithSendHeaders(Metadata headers) { + if (serverHeadersStartNanos > 0) { + long durationNanos = System.nanoTime() - serverHeadersStartNanos; + recordDuration(serverHeadersDuration, durationNanos); + serverHeadersStartNanos = 0; + } + super.sendHeaders(headers); + } + + @Override + public void sendMessage(InputStream message) { + if (dataPlaneCallClosed.get()) { + return; + } + + if (passThroughMode.get()) { + super.sendMessage(message); + return; + } + + try { + ByteString bodyByteString = outboundStreamToByteString(message); + ProcessingRequest requestToSend = null; + boolean sendRawImmediately = false; + + synchronized (streamLock) { + if (passThroughMode.get()) { + sendRawImmediately = true; + } else if (savedResponseHeaders != null) { + savedOutgoingMessages.add(new KnownLengthInputStream(bodyByteString)); + } else if (isExtProcStreamDraining() || isExtProcStreamCompleted()) { + pendingDrainingMessages.add(new KnownLengthInputStream(bodyByteString)); + } else if (currentProcessingMode.getResponseBodyMode() == ProcessingMode.BodySendMode.NONE) { + sendRawImmediately = true; + } else if (config.getObservabilityMode()) { + sendRawImmediately = true; + requestToSend = prepareResponseBodyRequest(bodyByteString); + } else { + // Flow control active + if (upstreamToSidestreamWindow <= 0 || !pendingResponseBodyMessages.isEmpty()) { + pendingResponseBodyMessages.add(bodyByteString); + } else { + upstreamToSidestreamWindow -= bodyByteString.size(); + requestToSend = prepareResponseBodyRequest(bodyByteString); + } + } + } + + if (sendRawImmediately) { + super.sendMessage(new KnownLengthInputStream(bodyByteString)); + if (requestToSend != null) { + sendToExtProc(requestToSend); + } + } else if (requestToSend != null) { + sendToExtProc(requestToSend); + } + } catch (IOException e) { + rawCall.close( + Status.INTERNAL.withDescription("Failed to serialize response body").withCause(e), + new Metadata()); + } + } + + @Override + public void close(Status status, Metadata trailers) { + serverTrailersStartNanos = System.nanoTime(); + if (isExtProcStreamFailed() + && !config.getObservabilityMode() + && (!config.getFailureModeAllow() || bodyMessageSentToExtProc.get())) { + if (markDataPlaneCallClosed(dataPlaneCallState)) { + proceedWithClose( + Status.INTERNAL.withDescription("External processor stream failed") + .withCause(status.getCause()), + new Metadata()); + } + return; + } + + synchronized (streamLock) { + if (passThroughMode.get()) { + if (markDataPlaneCallClosed(dataPlaneCallState)) { + proceedWithClose(status, trailers); + } + closeExtProcStream(); + return; + } + + this.savedStatus = status; + this.savedTrailers = trailers; + + if (!pendingResponseBodyMessages.isEmpty()) { + pendingClose.set(true); + return; + } + + if (isExtProcStreamCompleted()) { + proceedWithClose(); + return; + } + + if (savedResponseHeaders != null) { + return; + } + } + + if (!responseHeadersSent.get()) { + trailersOnly.set(true); + } + + triggerCloseHandshake(trailers); + + if (config.getObservabilityMode()) { + synchronized (streamLock) { + proceedWithClose(); + } + @SuppressWarnings("unused") + ScheduledFuture unused = scheduler.schedule( + this::closeExtProcStream, + config.getDeferredCloseTimeoutNanos(), + TimeUnit.NANOSECONDS); + } + } + + void proceedWithClose() { + synchronized (streamLock) { + if (savedStatus != null + && (isExtProcStreamCompleted() || config.getObservabilityMode())) { + if (markDataPlaneCallClosed(dataPlaneCallState)) { + proceedWithClose(savedStatus, savedTrailers); + } + savedStatus = null; + savedTrailers = null; + } + } + } + + private void proceedWithClose(Status status, Metadata trailers) { + if (dataPlaneCallClosed.compareAndSet(false, true)) { + + if (serverTrailersStartNanos > 0) { + long durationNanos = System.nanoTime() - serverTrailersStartNanos; + recordDuration(serverTrailersDuration, durationNanos); + serverTrailersStartNanos = 0; + } + super.close(status, trailers); + } + } + + private void triggerCloseHandshake(Metadata trailers) { + if (isExtProcStreamDraining()) { + return; + } + if (isExtProcStreamCompleted() || !terminationTriggered.compareAndSet(false, true)) { + return; + } + + boolean sendResponseHeaders = + currentProcessingMode.getResponseHeaderMode() == ProcessingMode.HeaderSendMode.SEND + || currentProcessingMode.getResponseHeaderMode() + == ProcessingMode.HeaderSendMode.DEFAULT; + + + boolean sendResponseTrailers = + currentProcessingMode.getResponseTrailerMode() == ProcessingMode.HeaderSendMode.SEND; + + if (trailersOnly.get()) { + if (sendResponseHeaders) { + sendToExtProc(ProcessingRequest.newBuilder() + .setResponseHeaders(HttpHeaders.newBuilder() + .setHeaders(toHeaderMap(trailers, config.getForwardRulesConfig())) + .setEndOfStream(true) + .build()) + .build()); + } else { + proceedWithClose(); + if (!config.getObservabilityMode()) { + closeExtProcStream(); + } + } + } else if (sendResponseTrailers) { + isProcessingTrailers.set(true); + sendToExtProc(ProcessingRequest.newBuilder() + .setResponseTrailers(HttpTrailers.newBuilder() + .setTrailers(toHeaderMap(trailers, config.getForwardRulesConfig())) + .build()) + .build()); + } else { + if (isRequestSideCompleted()) { + unblockAfterStreamComplete(); + closeExtProcStream(); + } + } + } + + @GuardedBy("streamLock") + private ProcessingRequest prepareResponseBodyRequest(ByteString body) { + if (isExtProcStreamCompleted() + || currentProcessingMode.getResponseBodyMode() != ProcessingMode.BodySendMode.GRPC) { + return null; + } + + HttpBody.Builder bodyBuilder = HttpBody.newBuilder() + .setBody(body) + .setEndOfStream(false); + bodyMessageSentToExtProc.set(true); + + ProcessingRequest.Builder builder = ProcessingRequest.newBuilder() + .setResponseBody(bodyBuilder.build()); + mergeAccumulatedWindowUpdates(builder); + return builder.build(); + } + + void drainPendingResponseBodyMessages() { + boolean triggerClose = false; + while (true) { + ProcessingRequest request = null; + synchronized (streamLock) { + if (upstreamToSidestreamWindow > 0 && !pendingResponseBodyMessages.isEmpty()) { + ByteString body = pendingResponseBodyMessages.poll(); + upstreamToSidestreamWindow -= body.size(); + request = prepareResponseBodyRequest(body); + } + if (request == null) { + if (pendingResponseBodyMessages.isEmpty() && pendingClose.get()) { + triggerClose = true; + pendingClose.set(false); + } + break; + } + } + if (request != null) { + sendToExtProc(request); + } + } + if (triggerClose) { + proceedWithClose(); + } + } + + private void handleRequestBodyResponse(BodyResponse bodyResponse) { + if (bodyResponse.hasResponse() && bodyResponse.getResponse().hasBodyMutation()) { + BodyMutation mutation = bodyResponse.getResponse().getBodyMutation(); + if (mutation.hasStreamedResponse()) { + StreamedBodyResponse streamed = mutation.getStreamedResponse(); + final int bodySize = streamed.getBody().size(); + synchronized (streamLock) { + sidestreamToUpstreamWindow -= bodySize; + } + deliverRequestBody(streamed); + } + } + } + + private void deliverRequestBody(StreamedBodyResponse streamed) { + synchronized (streamLock) { + pendingMutatedRequestBodies.add(streamed); + } + drainPendingMutatedRequestBodies(); + } + + int drainPendingMutatedRequestBodies() { + List toDeliver = new ArrayList<>(); + synchronized (streamLock) { + while (pendingAppRequests.get() > 0 && !pendingMutatedRequestBodies.isEmpty()) { + StreamedBodyResponse streamed = pendingMutatedRequestBodies.poll(); + pendingAppRequests.decrementAndGet(); + toDeliver.add(streamed); + } + } + for (StreamedBodyResponse streamed : toDeliver) { + final StreamedBodyResponse finalStreamed = streamed; + final int bodySize = streamed.getBody().size(); + callContext.run(() -> { + try { + if (!finalStreamed.getEndOfStreamWithoutMessage()) { + wrappedListener.onExternalBody(finalStreamed.getBody()); + } + if (finalStreamed.getEndOfStream() || finalStreamed.getEndOfStreamWithoutMessage()) { + wrappedListener.proceedWithHalfClose(); + } + } finally { + synchronized (streamLock) { + accumulatedWindowUpdateSidestreamToUpstream += bodySize; + } + trySendAccumulatedWindowUpdates(); + } + }); + } + return toDeliver.size(); + } + + private void handleResponseBodyResponse(BodyResponse bodyResponse) { + if (dataPlaneCallClosed.get()) { + return; + } + if (bodyResponse.hasResponse() && bodyResponse.getResponse().hasBodyMutation()) { + BodyMutation mutation = bodyResponse.getResponse().getBodyMutation(); + if (mutation.hasStreamedResponse()) { + StreamedBodyResponse streamed = mutation.getStreamedResponse(); + ByteString body = streamed.getBody(); + final int bodySize = body.size(); + synchronized (streamLock) { + sidestreamToDownstreamWindow -= bodySize; + } + deliverResponseBodyToClient(body); + } + } + } + + private void deliverResponseBodyToClient(ByteString body) { + boolean shouldSend = false; + synchronized (streamLock) { + if (super.isReady() && pendingDownstreamBodyMessages.isEmpty()) { + shouldSend = true; + } else { + pendingDownstreamBodyMessages.add(body); + } + } + if (shouldSend) { + final int bodySize = body.size(); + super.sendMessage(new KnownLengthInputStream(body)); + synchronized (streamLock) { + accumulatedWindowUpdateSidestreamToDownstream += bodySize; + } + trySendAccumulatedWindowUpdates(); + } + } + + void drainPendingDownstreamBodyMessages() { + while (true) { + ByteString body = null; + synchronized (streamLock) { + if (super.isReady() && !pendingDownstreamBodyMessages.isEmpty()) { + body = pendingDownstreamBodyMessages.poll(); + } + if (body == null) { + break; + } + } + if (body != null) { + final int bodySize = body.size(); + super.sendMessage(new KnownLengthInputStream(body)); + synchronized (streamLock) { + accumulatedWindowUpdateSidestreamToDownstream += bodySize; + } + trySendAccumulatedWindowUpdates(); + } + } + } + + private void handleImmediateResponse(ImmediateResponse immediate) + throws HeaderMutationDisallowedException { + Status status = Status.fromCodeValue(immediate.getGrpcStatus().getStatus()); + if (!immediate.getDetails().isEmpty()) { + status = status.withDescription(immediate.getDetails()); + } + + Metadata trailers = new Metadata(); + if (immediate.hasHeaders()) { + applyHeaderMutations(trailers, immediate.getHeaders(), mutationFilter, mutator); + } + + synchronized (streamLock) { + savedStatus = status; + savedTrailers = trailers; + } + + if (isProcessingTrailers.get()) { + unblockAfterStreamComplete(); + } else { + proceedWithClose(status, trailers); + unblockAfterStreamComplete(); + } + closeExtProcStream(); + } + + private void drainPendingDrainingMessages() { + synchronized (streamLock) { + InputStream msg; + while ((msg = pendingDrainingMessages.poll()) != null) { + super.sendMessage(msg); + } + passThroughMode.set(true); + } + } + + private void drainRequestMessagesFailOpen() { + List mutatedToDeliver = new ArrayList<>(); + synchronized (streamLock) { + StreamedBodyResponse streamed; + while ((streamed = pendingMutatedRequestBodies.poll()) != null) { + mutatedToDeliver.add(streamed); + } + } + for (StreamedBodyResponse streamed : mutatedToDeliver) { + final StreamedBodyResponse finalStreamed = streamed; + callContext.run(() -> { + if (!finalStreamed.getEndOfStreamWithoutMessage()) { + wrappedListener.onExternalBody(finalStreamed.getBody()); + } + if (finalStreamed.getEndOfStream() || finalStreamed.getEndOfStreamWithoutMessage()) { + wrappedListener.proceedWithHalfClose(); + } + }); + } + + List rawToDeliver = new ArrayList<>(); + synchronized (streamLock) { + ByteString body; + while ((body = pendingRequestBodyMessages.poll()) != null) { + rawToDeliver.add(body); + } + } + for (ByteString body : rawToDeliver) { + final ByteString finalBody = body; + callContext.run(() -> wrappedListener.onExternalBody(finalBody)); + } + + wrappedListener.drainSavedMessages(); + } + + void drainResponseMessagesFailOpen() { + boolean triggerClose = false; + while (true) { + Object msg = null; + boolean isByteString = false; + + synchronized (streamLock) { + if (super.isReady() && !pendingDownstreamBodyMessages.isEmpty()) { + msg = pendingDownstreamBodyMessages.poll(); + isByteString = true; + } else if (super.isReady() && pendingDownstreamBodyMessages.isEmpty() + && !pendingResponseBodyMessages.isEmpty()) { + msg = pendingResponseBodyMessages.poll(); + isByteString = true; + } else if (super.isReady() && pendingDownstreamBodyMessages.isEmpty() + && pendingResponseBodyMessages.isEmpty() + && !savedOutgoingMessages.isEmpty()) { + msg = savedOutgoingMessages.poll(); + isByteString = false; + } else if (super.isReady() && pendingDownstreamBodyMessages.isEmpty() + && pendingResponseBodyMessages.isEmpty() + && savedOutgoingMessages.isEmpty() + && !pendingDrainingMessages.isEmpty()) { + msg = pendingDrainingMessages.poll(); + isByteString = false; + } + + if (msg == null) { + if (pendingDownstreamBodyMessages.isEmpty() + && pendingResponseBodyMessages.isEmpty() + && savedOutgoingMessages.isEmpty() + && pendingDrainingMessages.isEmpty()) { + passThroughMode.set(true); + if (pendingClose.get()) { + triggerClose = true; + pendingClose.set(false); + } + } + break; + } + } + + if (msg != null) { + if (isByteString) { + super.sendMessage(new KnownLengthInputStream((ByteString) msg)); + } else { + super.sendMessage((InputStream) msg); + } + } + } + if (triggerClose) { + proceedWithClose(); + } + } + + private void handleFailOpen() { + activateCall(); + drainRequestMessagesFailOpen(); + proceedWithSendHeaders(); + drainResponseMessagesFailOpen(); + closeExtProcStream(); + wrappedListener.onReadyNotify(); + } + + /** + * Evaluates whether the external processor stream can be safely closed and the + * data plane call terminated. + * + *

This method acts as a cleanup checkpoint. It is invoked when request-side + * processing completes (e.g., half-close) or when call termination is triggered. + * + *

The stream is only closed if: + *

    + *
  • Call termination has been initiated ({@code terminationTriggered} is true).
  • + *
  • The request side of the call is fully completed ({@code isRequestSideCompleted} + * is true).
  • + *
  • There are no outstanding response-side messages (such as mutated response headers + * or trailers) expected from the external processor.
  • + *
+ * + *

If all conditions are met, the data plane call is unblocked to allow the close status + * and trailers to be propagated, and the external processor gRPC stream is terminated. + */ + private void checkEndOfStream() { + if (terminationTriggered.get() && isRequestSideCompleted() + && !expectedResponses.contains(EventType.RESPONSE_HEADERS) + && !expectedResponses.contains(EventType.RESPONSE_TRAILERS)) { + unblockAfterStreamComplete(); + closeExtProcStream(); + } + } + + private boolean isRequestSideCompleted() { + return (currentProcessingMode.getRequestHeaderMode() != ProcessingMode.HeaderSendMode.SEND + && currentProcessingMode.getRequestBodyMode() != ProcessingMode.BodySendMode.GRPC) + || requestSideClosed.get(); + } + + void unblockAfterStreamComplete() { + proceedWithSendHeaders(); + drainPendingDrainingMessages(); + wrappedListener.drainSavedMessages(); + wrappedListener.onReadyNotify(); + proceedWithClose(); + } + } + + static final class DataPlaneServerListener extends ServerCall.Listener { + private final DataPlaneServerCall dataPlaneServerCall; + final Queue savedMessages = new ConcurrentLinkedQueue<>(); + private volatile boolean halfCloseReceived; + private volatile boolean halfCloseDeferred; + private volatile ServerCall.Listener delegate; + + private DataPlaneServerListener(DataPlaneServerCall dataPlaneServerCall) { + this.dataPlaneServerCall = dataPlaneServerCall; + } + + void setDelegate(ServerCall.Listener delegate) { + dataPlaneServerCall.triggerEvent(new SetDelegateEvent(delegate)); + } + + private void handleSetDelegate(ServerCall.Listener delegate) { + this.delegate = delegate; + dataPlaneServerCall.callContext.run(() -> { + InputStream msg; + while ((msg = savedMessages.poll()) != null) { + delegate.onMessage(msg); + } + if (halfCloseReceived) { + proceedWithHalfClose(); + } + }); + } + + @Override + public void onEvent(Object event) { + if (dataPlaneServerCall.dataPlaneCallClosed.get()) { + return; + } + + if (event instanceof ExtProcResponseEvent) { + dataPlaneServerCall.handleExtProcResponse(((ExtProcResponseEvent) event).getResponse()); + } else if (event instanceof ExtProcErrorEvent) { + dataPlaneServerCall.handleExtProcError(((ExtProcErrorEvent) event).getCause()); + } else if (event instanceof ExtProcCompletedEvent) { + dataPlaneServerCall.handleExtProcCompleted(); + } else if (event instanceof ExtProcStreamReadyEvent) { + dataPlaneServerCall.onExtProcStreamReady(); + } else if (event instanceof SetDelegateEvent) { + handleSetDelegate(((SetDelegateEvent) event).getDelegate()); + } + } + + void drainSavedMessages() { + ServerCall.Listener del = delegate; + if (del != null) { + dataPlaneServerCall.callContext.run(() -> { + InputStream msg; + while ((msg = savedMessages.poll()) != null) { + del.onMessage(msg); + } + if (halfCloseReceived) { + proceedWithHalfClose(); + } + }); + } + } + + @Override + public void onReady() { + if (dataPlaneServerCall.passThroughMode.get()) { + onReadyNotify(); + return; + } + if (dataPlaneServerCall.isExtProcStreamCompleted()) { + dataPlaneServerCall.drainResponseMessagesFailOpen(); + return; + } + dataPlaneServerCall.drainPendingDownstreamBodyMessages(); + dataPlaneServerCall.drainPendingRequests(); + onReadyNotify(); + } + + void onReadyNotify() { + ServerCall.Listener del = delegate; + if (del != null && dataPlaneServerCall.isReady()) { + dataPlaneServerCall.callContext.run(del::onReady); + } + } + + @Override + public void onMessage(InputStream message) { + if (dataPlaneServerCall.dataPlaneCallClosed.get()) { + return; + } + if (dataPlaneServerCall.requestSideClosed.get()) { + return; + } + ServerCall.Listener del = delegate; + if (dataPlaneServerCall.passThroughMode.get() && del != null) { + dataPlaneServerCall.callContext.run(() -> del.onMessage(message)); + return; + } + + if (dataPlaneServerCall.isExtProcStreamCompleted() + || dataPlaneServerCall.isExtProcStreamDraining() + || dataPlaneServerCall.currentProcessingMode.getRequestBodyMode() + != ProcessingMode.BodySendMode.GRPC + || dataPlaneServerCall.config.getObservabilityMode()) { + + if (del == null || dataPlaneServerCall.isExtProcStreamDraining()) { + try { + ByteString copiedBytes = ByteString.readFrom(message); + savedMessages.add(new KnownLengthInputStream(copiedBytes)); + } catch (IOException e) { + dataPlaneServerCall.rawCall.close( + Status.INTERNAL.withDescription("Failed to buffer client request").withCause(e), + new Metadata()); + } + } else { + dataPlaneServerCall.callContext.run(() -> del.onMessage(message)); + } + return; + } + + // Flow control active + try { + ByteString bodyByteString = ByteString.readFrom(message); + synchronized (dataPlaneServerCall.streamLock) { + // Re-check stream state under lock + if (dataPlaneServerCall.isExtProcStreamCompleted() + || dataPlaneServerCall.isExtProcStreamDraining()) { + if (del == null || dataPlaneServerCall.isExtProcStreamDraining()) { + savedMessages.add(new KnownLengthInputStream(bodyByteString)); + } else { + dataPlaneServerCall.callContext.run( + () -> del.onMessage(new KnownLengthInputStream(bodyByteString))); + } + return; + } + + if (dataPlaneServerCall.downstreamToSidestreamWindow <= 0 + || !dataPlaneServerCall.pendingRequestBodyMessages.isEmpty()) { + dataPlaneServerCall.pendingRequestBodyMessages.add(bodyByteString); + } else { + sendRequestBodyToExtProc(bodyByteString); + } + } + dataPlaneServerCall.drainPendingRequests(); + } catch (IOException e) { + dataPlaneServerCall.rawCall.close( + Status.INTERNAL.withDescription("Failed to read client request").withCause(e), + new Metadata()); + } + } + + @Override + public void onHalfClose() { + if (dataPlaneServerCall.dataPlaneCallClosed.get()) { + return; + } + if (dataPlaneServerCall.requestSideClosed.get()) { + return; + } + dataPlaneServerCall.clientHalfCloseStartNanos = System.nanoTime(); + dataPlaneServerCall.halfClosed.set(true); + halfCloseReceived = true; + if (dataPlaneServerCall.isExtProcStreamDraining()) { + return; + } + ServerCall.Listener del = delegate; + if ((dataPlaneServerCall.passThroughMode.get() + || dataPlaneServerCall.isExtProcStreamCompleted()) && del != null) { + proceedWithHalfClose(); + return; + } + + if (dataPlaneServerCall.dataPlaneCallState.get() == DataPlaneCallState.IDLE) { + halfCloseDeferred = true; + return; + } + + if (dataPlaneServerCall.currentProcessingMode.getRequestBodyMode() + == ProcessingMode.BodySendMode.NONE) { + proceedWithHalfClose(); + return; + } + + synchronized (dataPlaneServerCall.streamLock) { + if (!dataPlaneServerCall.pendingRequestBodyMessages.isEmpty()) { + halfCloseDeferred = true; + } else { + sendHalfCloseToExtProc(); + } + } + } + + void handleDeferredHalfClose() { + if (dataPlaneServerCall.currentProcessingMode.getRequestBodyMode() + == ProcessingMode.BodySendMode.NONE + || dataPlaneServerCall.isExtProcStreamCompleted()) { + proceedWithHalfClose(); + } else { + synchronized (dataPlaneServerCall.streamLock) { + sendHalfCloseToExtProc(); + } + } + } + + void proceedWithHalfClose() { + if (!dataPlaneServerCall.requestSideClosed.compareAndSet(false, true)) { + return; + } + halfCloseReceived = true; + if (dataPlaneServerCall.clientHalfCloseStartNanos > 0) { + long durationNanos = System.nanoTime() - dataPlaneServerCall.clientHalfCloseStartNanos; + dataPlaneServerCall.recordDuration(clientHalfCloseDuration, durationNanos); + dataPlaneServerCall.clientHalfCloseStartNanos = 0; + } + ServerCall.Listener del = delegate; + if (del != null) { + dataPlaneServerCall.callContext.run(del::onHalfClose); + } + dataPlaneServerCall.checkEndOfStream(); + } + + void onExternalBody(ByteString body) { + ServerCall.Listener del = delegate; + // In the future, if zero-copy reads are needed downstream, this can be optimized + // by wrapping the ByteString in an InputStream that implements HasByteBuffer, + // KnownLength, and Detachable. + if (del != null) { + dataPlaneServerCall.callContext.run(() -> del.onMessage(body.newInput())); + } else { + savedMessages.add(body.newInput()); + } + } + + @GuardedBy("dataPlaneServerCall.streamLock") + private void sendRequestBodyToExtProc(ByteString bodyByteString) { + if (dataPlaneServerCall.isExtProcStreamCompleted() + || dataPlaneServerCall.currentProcessingMode.getRequestBodyMode() + != ProcessingMode.BodySendMode.GRPC) { + return; + } + + dataPlaneServerCall.downstreamToSidestreamWindow -= bodyByteString.size(); + dataPlaneServerCall.bodyMessageSentToExtProc.set(true); + + HttpBody.Builder bodyBuilder = HttpBody.newBuilder() + .setBody(bodyByteString) + .setEndOfStream(false); + + ProcessingRequest.Builder builder = ProcessingRequest.newBuilder() + .setRequestBody(bodyBuilder.build()); + dataPlaneServerCall.mergeAccumulatedWindowUpdates(builder); + dataPlaneServerCall.sendToExtProc(builder.build()); + } + + @GuardedBy("dataPlaneServerCall.streamLock") + private void sendHalfCloseToExtProc() { + if (dataPlaneServerCall.isExtProcStreamCompleted() + || dataPlaneServerCall.currentProcessingMode.getRequestBodyMode() + != ProcessingMode.BodySendMode.GRPC) { + return; + } + + HttpBody.Builder bodyBuilder = HttpBody.newBuilder() + .setEndOfStreamWithoutMessage(true); + + ProcessingRequest.Builder builder = ProcessingRequest.newBuilder() + .setRequestBody(bodyBuilder.build()); + dataPlaneServerCall.mergeAccumulatedWindowUpdates(builder); + dataPlaneServerCall.sendToExtProc(builder.build()); + } + + void drainPendingRequestBodyMessages() { + boolean triggerHalfClose = false; + while (true) { + ProcessingRequest request = null; + synchronized (dataPlaneServerCall.streamLock) { + if (dataPlaneServerCall.downstreamToSidestreamWindow > 0 + && !dataPlaneServerCall.pendingRequestBodyMessages.isEmpty()) { + ByteString body = dataPlaneServerCall.pendingRequestBodyMessages.poll(); + dataPlaneServerCall.downstreamToSidestreamWindow -= body.size(); + dataPlaneServerCall.bodyMessageSentToExtProc.set(true); + + HttpBody.Builder bodyBuilder = HttpBody.newBuilder() + .setBody(body) + .setEndOfStream(false); + ProcessingRequest.Builder builder = ProcessingRequest.newBuilder() + .setRequestBody(bodyBuilder.build()); + dataPlaneServerCall.mergeAccumulatedWindowUpdates(builder); + request = builder.build(); + } + + if (request == null) { + if (dataPlaneServerCall.pendingRequestBodyMessages.isEmpty() + && halfCloseDeferred) { + triggerHalfClose = true; + halfCloseDeferred = false; + } + break; + } + } + + if (request != null) { + dataPlaneServerCall.sendToExtProc(request); + } + } + + if (triggerHalfClose) { + synchronized (dataPlaneServerCall.streamLock) { + sendHalfCloseToExtProc(); + } + } + } + + @Override + public void onCancel() { + dataPlaneServerCall.cancelExtProcStream( + Status.CANCELLED.withDescription("Client cancelled RPC").asRuntimeException()); + ServerCall.Listener del = delegate; + if (del != null) { + dataPlaneServerCall.callContext.run(del::onCancel); + } + } + + @Override + public void onComplete() { + ServerCall.Listener del = delegate; + if (del != null) { + dataPlaneServerCall.callContext.run(del::onComplete); + } + } + } + + static final class ExtProcResponseEvent { + private final ProcessingResponse response; + + ExtProcResponseEvent(ProcessingResponse response) { + this.response = response; + } + + ProcessingResponse getResponse() { + return response; + } + } + + static final class ExtProcErrorEvent { + private final Throwable cause; + + ExtProcErrorEvent(Throwable cause) { + this.cause = cause; + } + + Throwable getCause() { + return cause; + } + } + + static final class ExtProcCompletedEvent {} + + static final class ExtProcStreamReadyEvent {} + + static final class SetDelegateEvent { + private final ServerCall.Listener delegate; + + SetDelegateEvent(ServerCall.Listener delegate) { + this.delegate = delegate; + } + + ServerCall.Listener getDelegate() { + return delegate; + } + } +} diff --git a/xds/src/test/java/io/grpc/xds/ExternalProcessorServerInterceptorTest.java b/xds/src/test/java/io/grpc/xds/ExternalProcessorServerInterceptorTest.java new file mode 100644 index 00000000000..eb5480a6a9d --- /dev/null +++ b/xds/src/test/java/io/grpc/xds/ExternalProcessorServerInterceptorTest.java @@ -0,0 +1,13188 @@ +/* + * Copyright 2024 The gRPC Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package io.grpc.xds; + +import static com.google.common.truth.Truth.assertThat; +import static com.google.common.truth.Truth.assertWithMessage; + +import com.google.common.io.ByteStreams; +import com.google.protobuf.Any; +import com.google.protobuf.ByteString; +import io.envoyproxy.envoy.config.core.v3.GrpcService; +import io.envoyproxy.envoy.extensions.filters.http.ext_proc.v3.ExternalProcessor; +import io.envoyproxy.envoy.extensions.filters.http.ext_proc.v3.HeaderForwardingRules; +import io.envoyproxy.envoy.extensions.filters.http.ext_proc.v3.ProcessingMode; +import io.envoyproxy.envoy.service.ext_proc.v3.BodyMutation; +import io.envoyproxy.envoy.service.ext_proc.v3.BodyResponse; +import io.envoyproxy.envoy.service.ext_proc.v3.CommonResponse; +import io.envoyproxy.envoy.service.ext_proc.v3.ExternalProcessorGrpc; +import io.envoyproxy.envoy.service.ext_proc.v3.HeaderMutation; +import io.envoyproxy.envoy.service.ext_proc.v3.HeadersResponse; +import io.envoyproxy.envoy.service.ext_proc.v3.ProcessingRequest; +import io.envoyproxy.envoy.service.ext_proc.v3.ProcessingResponse; +import io.envoyproxy.envoy.service.ext_proc.v3.StreamedBodyResponse; +import io.envoyproxy.envoy.service.ext_proc.v3.TrailersResponse; +import io.grpc.ClientCall; +import io.grpc.Context; +import io.grpc.Contexts; +import io.grpc.ForwardingServerCall; +import io.grpc.Metadata; +import io.grpc.MethodDescriptor; +import io.grpc.NameResolver; +import io.grpc.NameResolverProvider; +import io.grpc.NameResolverRegistry; +import io.grpc.ServerCall; +import io.grpc.ServerCallHandler; +import io.grpc.ServerInterceptor; +import io.grpc.ServerInterceptors; +import io.grpc.ServerServiceDefinition; +import io.grpc.Status; +import io.grpc.inprocess.InProcessChannelBuilder; +import io.grpc.inprocess.InProcessServerBuilder; +import io.grpc.stub.ServerCallStreamObserver; +import io.grpc.stub.StreamObserver; +import io.grpc.testing.GrpcCleanupRule; +import io.grpc.util.MutableHandlerRegistry; +import io.grpc.xds.ExternalProcessorFilter.ExternalProcessorFilterConfig; +import io.grpc.xds.ExternalProcessorFilter.ExternalProcessorFilterOverrideConfig; +import io.grpc.xds.Filter.FilterContext; +import io.grpc.xds.client.Bootstrapper; +import io.grpc.xds.client.EnvoyProtoData.Node; +import io.grpc.xds.internal.grpcservice.CachedChannelManager; +import java.io.ByteArrayInputStream; +import java.io.IOException; +import java.io.InputStream; +import java.net.SocketAddress; +import java.net.URI; +import java.nio.charset.StandardCharsets; +import java.util.ArrayList; +import java.util.Collection; +import java.util.Collections; +import java.util.List; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; +import org.junit.Before; +import org.junit.Rule; +import org.junit.Test; +import org.junit.runner.RunWith; +import org.junit.runners.JUnit4; +import org.mockito.Mockito; + +@RunWith(JUnit4.class) +public class ExternalProcessorServerInterceptorTest { + private static final Context.Key TRACE_KEY = Context.key("trace-id"); + + @Rule public final GrpcCleanupRule grpcCleanup = new GrpcCleanupRule(); + + private String extProcServerName; + private ExternalProcessorFilter.Provider provider; + private static final FilterContext FAKE_CONTEXT = + FilterContext.create("test-filter", new io.grpc.MetricRecorder() {}); + private Filter.FilterConfigParseContext filterContext; + private Bootstrapper.BootstrapInfo bootstrapInfo; + private Bootstrapper.ServerInfo serverInfo; + + private static class InputStreamMarshaller implements MethodDescriptor.Marshaller { + @Override + public InputStream stream(InputStream value) { + try { + byte[] bytes = ByteStreams.toByteArray(value); + return new ByteArrayInputStream(bytes); + } catch (IOException e) { + throw new RuntimeException(e); + } + } + + @Override + public InputStream parse(InputStream stream) { + try { + byte[] bytes = ByteStreams.toByteArray(stream); + return new ByteArrayInputStream(bytes); + } catch (IOException e) { + throw new RuntimeException(e); + } + } + } + + private static final MethodDescriptor METHOD_SAY_HELLO_RAW = + MethodDescriptor.newBuilder() + .setType(MethodDescriptor.MethodType.UNARY) + .setFullMethodName("test.TestService/SayHello") + .setRequestMarshaller(new InputStreamMarshaller()) + .setResponseMarshaller(new InputStreamMarshaller()) + .build(); + + private static final MethodDescriptor METHOD_SAY_HELLO_BIDI = + MethodDescriptor.newBuilder() + .setType(MethodDescriptor.MethodType.BIDI_STREAMING) + .setFullMethodName("test.TestService/SayHelloBidi") + .setRequestMarshaller(new InputStreamMarshaller()) + .setResponseMarshaller(new InputStreamMarshaller()) + .build(); + + private static final MethodDescriptor + METHOD_SAY_HELLO_CLIENT_STREAMING = + MethodDescriptor.newBuilder() + .setType(MethodDescriptor.MethodType.CLIENT_STREAMING) + .setFullMethodName("test.TestService/SayHelloClientStreaming") + .setRequestMarshaller(new InputStreamMarshaller()) + .setResponseMarshaller(new InputStreamMarshaller()) + .build(); + + private String dataPlaneServerName; + private io.grpc.Channel dataPlaneChannel; + + private interface DataPlaneServiceHandler { + default void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + + default StreamObserver sayHelloBidi(StreamObserver responseObserver) { + return new StreamObserver() { + @Override + public void onNext(InputStream value) { + responseObserver.onNext(value); + } + + @Override + public void onError(Throwable t) { + responseObserver.onError(t); + } + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + + default StreamObserver sayHelloClientStreaming( + StreamObserver responseObserver) { + return new StreamObserver() { + private InputStream lastValue; + + @Override + public void onNext(InputStream value) { + lastValue = value; + } + + @Override + public void onError(Throwable t) { + responseObserver.onError(t); + } + + @Override + public void onCompleted() { + responseObserver.onNext(lastValue); + responseObserver.onCompleted(); + } + }; + } + } + + private volatile DataPlaneServiceHandler dataPlaneHandler = new DataPlaneServiceHandler() {}; + + private void startDataPlane(ServerInterceptor... interceptors) throws Exception { + dataPlaneServerName = InProcessServerBuilder.generateName(); + ServerServiceDefinition dataPlaneService = + ServerServiceDefinition.builder("test.TestService") + .addMethod( + METHOD_SAY_HELLO_RAW, + io.grpc.stub.ServerCalls.asyncUnaryCall( + (request, responseObserver) -> + dataPlaneHandler.sayHello(request, responseObserver))) + .addMethod( + METHOD_SAY_HELLO_BIDI, + io.grpc.stub.ServerCalls.asyncBidiStreamingCall( + responseObserver -> dataPlaneHandler.sayHelloBidi(responseObserver))) + .addMethod( + METHOD_SAY_HELLO_CLIENT_STREAMING, + io.grpc.stub.ServerCalls.asyncClientStreamingCall( + responseObserver -> dataPlaneHandler.sayHelloClientStreaming(responseObserver))) + .build(); + + grpcCleanup.register( + InProcessServerBuilder.forName(dataPlaneServerName) + .addService( + ServerInterceptors.intercept( + dataPlaneService, java.util.Arrays.asList(interceptors))) + .directExecutor() + .build() + .start()); + + dataPlaneChannel = + grpcCleanup.register( + InProcessChannelBuilder.forName(dataPlaneServerName).directExecutor().build()); + } + + private static class SimpleServerCall extends ServerCall { + private final MethodDescriptor method; + + SimpleServerCall(MethodDescriptor method) { + this.method = method; + } + + @Override + public void request(int numMessages) {} + + @Override + public void sendHeaders(Metadata headers) {} + + @Override + public void sendMessage(InputStream message) {} + + @Override + public void close(Status status, Metadata trailers) {} + + @Override + public boolean isCancelled() { + return false; + } + + @Override + public MethodDescriptor getMethodDescriptor() { + return method; + } + } + + private static class InProcessNameResolverProvider extends NameResolverProvider { + @Override + public NameResolver newNameResolver(URI targetUri, NameResolver.Args args) { + if ("in-process".equals(targetUri.getScheme())) { + return new NameResolver() { + @Override + public String getServiceAuthority() { + return "localhost"; + } + + @Override + public void start(Listener2 listener) {} + + @Override + public void shutdown() {} + }; + } + return null; + } + + @Override + protected boolean isAvailable() { + return true; + } + + @Override + protected int priority() { + return 5; + } + + @Override + public String getDefaultScheme() { + return "in-process"; + } + + @Override + public Collection> getProducedSocketAddressTypes() { + return Collections.emptyList(); + } + } + + @Before + public void setUp() throws Exception { + NameResolverRegistry.getDefaultRegistry().register(new InProcessNameResolverProvider()); + + extProcServerName = InProcessServerBuilder.generateName(); + provider = new ExternalProcessorFilter.Provider(); + + bootstrapInfo = + Bootstrapper.BootstrapInfo.builder() + .node(Node.newBuilder().build()) + .servers( + Collections.singletonList( + Bootstrapper.ServerInfo.create("test_target", Collections.emptyMap()))) + .build(); + + serverInfo = + Bootstrapper.ServerInfo.create( + "test_target", Collections.emptyMap(), false, true, false, false, null); + + filterContext = + Filter.FilterConfigParseContext.builder() + .bootstrapInfo(bootstrapInfo) + .serverInfo(serverInfo) + .build(); + } + + private ExternalProcessor.Builder createBaseProto(String targetName) { + return ExternalProcessor.newBuilder() + .setGrpcService( + GrpcService.newBuilder() + .setGoogleGrpc( + GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///" + targetName) + .addChannelCredentialsPlugin( + Any.newBuilder() + .setTypeUrl( + "type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build()); + } + + // ============================================================================ + // Category 1: Configuration Override + // ============================================================================ + + @Test + public void givenOverrideConfig_whenGrpcServiceOverridden_thenUsesNewService() throws Exception { + ExternalProcessor parentProto = + createBaseProto(extProcServerName) + .setGrpcService( + GrpcService.newBuilder() + .setGoogleGrpc( + GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///parent") + .addChannelCredentialsPlugin( + Any.newBuilder() + .setTypeUrl( + "type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build()) + .build(); + + GrpcService overrideService = + GrpcService.newBuilder() + .setGoogleGrpc( + GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///override") + .addChannelCredentialsPlugin( + Any.newBuilder() + .setTypeUrl( + "type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build(); + io.envoyproxy.envoy.extensions.filters.http.ext_proc.v3.ExtProcPerRoute perRoute = + io.envoyproxy.envoy.extensions.filters.http.ext_proc.v3.ExtProcPerRoute.newBuilder() + .setOverrides( + io.envoyproxy.envoy.extensions.filters.http.ext_proc.v3.ExtProcOverrides + .newBuilder() + .setGrpcService(overrideService) + .build()) + .build(); + + ConfigOrError parentResult = + provider.parseFilterConfig(Any.pack(parentProto), filterContext); + assertThat(parentResult.errorDetail).isNull(); + ExternalProcessorFilterConfig parentConfig = parentResult.config; + ConfigOrError overrideResult = + provider.parseFilterConfigOverride(Any.pack(perRoute), filterContext); + assertThat(overrideResult.errorDetail).isNull(); + ExternalProcessorFilterOverrideConfig overrideConfig = overrideResult.config; + + ExternalProcessorFilter filter = new ExternalProcessorFilter(FAKE_CONTEXT); + ExternalProcessorServerInterceptor interceptor = + (ExternalProcessorServerInterceptor) + filter.buildServerInterceptor(parentConfig, overrideConfig); + + assertThat( + interceptor + .getFilterConfig() + .getExternalProcessor() + .getGrpcService() + .getGoogleGrpc() + .getTargetUri()) + .isEqualTo("in-process:///override"); + } + + // ============================================================================ + // Category 2: Server Interceptor & Lifecycle + // ============================================================================ + + @Test + @SuppressWarnings("unchecked") + public void givenInterceptor_whenCallIntercepted_thenExtProcStubUsesSynchronizationContext() + throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + ExternalProcessor proto = + ExternalProcessor.newBuilder() + .setGrpcService( + GrpcService.newBuilder() + .setGoogleGrpc( + GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///" + uniqueExtProcServerName) + .addChannelCredentialsPlugin( + Any.newBuilder() + .setTypeUrl( + "type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + // External Processor Server + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + responseObserverRef.set(responseObserver); + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) {} + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + final AtomicReference capturedExecutor = new AtomicReference<>(); + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName) + .directExecutor() + .intercept( + new io.grpc.ClientInterceptor() { + @Override + public io.grpc.ClientCall interceptCall( + MethodDescriptor method, + io.grpc.CallOptions callOptions, + io.grpc.Channel next) { + if (method.equals(ExternalProcessorGrpc.getProcessMethod())) { + capturedExecutor.set(callOptions.getExecutor()); + } + return next.newCall(method, callOptions); + } + }) + .build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + MutableHandlerRegistry uniqueRegistry = new MutableHandlerRegistry(); + String uniqueDataPlaneServerName = InProcessServerBuilder.generateName(); + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueDataPlaneServerName) + .fallbackHandlerRegistry(uniqueRegistry) + .directExecutor() + .build() + .start()); + + io.grpc.ManagedChannel dataPlaneChannel = + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueDataPlaneServerName).directExecutor().build()); + + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + ServerCallHandler nextHandler = + (call, headers) -> { + return new ServerCall.Listener() {}; + }; + + try { + interceptor.interceptCall( + new SimpleServerCall(METHOD_SAY_HELLO_RAW), new Metadata(), nextHandler); + + assertThat(capturedExecutor.get()).isNotNull(); + assertThat(capturedExecutor.get()).isSameInstanceAs(com.google.common.util.concurrent.MoreExecutors.directExecutor()); + } finally { + if (responseObserverRef.get() != null) { + responseObserverRef.get().onCompleted(); + } + channelManager.close(); + } + } + + @Test + @SuppressWarnings("unchecked") + public void givenGrpcServiceWithTimeout_whenCallIntercepted_thenExtProcStubHasCorrectDeadline() + throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + ExternalProcessor proto = + ExternalProcessor.newBuilder() + .setGrpcService( + GrpcService.newBuilder() + .setGoogleGrpc( + GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///" + uniqueExtProcServerName) + .addChannelCredentialsPlugin( + Any.newBuilder() + .setTypeUrl( + "type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .setTimeout(com.google.protobuf.Duration.newBuilder().setSeconds(5).build()) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + // External Processor Server + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + responseObserverRef.set(responseObserver); + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) {} + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + final AtomicReference capturedDeadline = new AtomicReference<>(); + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName) + .directExecutor() + .intercept( + new io.grpc.ClientInterceptor() { + @Override + public io.grpc.ClientCall interceptCall( + MethodDescriptor method, + io.grpc.CallOptions callOptions, + io.grpc.Channel next) { + if (method.equals(ExternalProcessorGrpc.getProcessMethod())) { + capturedDeadline.set(callOptions.getDeadline()); + } + return next.newCall(method, callOptions); + } + }) + .build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + MutableHandlerRegistry uniqueRegistry = new MutableHandlerRegistry(); + String uniqueDataPlaneServerName = InProcessServerBuilder.generateName(); + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueDataPlaneServerName) + .fallbackHandlerRegistry(uniqueRegistry) + .directExecutor() + .build() + .start()); + + io.grpc.ManagedChannel dataPlaneChannel = + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueDataPlaneServerName).directExecutor().build()); + + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + ServerCallHandler nextHandler = + (call, headers) -> { + return new ServerCall.Listener() {}; + }; + + try { + interceptor.interceptCall( + new SimpleServerCall(METHOD_SAY_HELLO_RAW), new Metadata(), nextHandler); + + assertThat(capturedDeadline.get()).isNotNull(); + assertThat(capturedDeadline.get().timeRemaining(TimeUnit.SECONDS)).isAtLeast(4); + } finally { + if (responseObserverRef.get() != null) { + responseObserverRef.get().onCompleted(); + } + channelManager.close(); + } + } + + // ============================================================================ + // Category 3: Protocol config propagation + // ============================================================================ + + @Test + public void protocolConfig_onHeaders() throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + + final CountDownLatch extProcLatch = + new CountDownLatch(2); // Expecting request headers and request body + final List capturedRequests = + Collections.synchronizedList(new ArrayList<>()); + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + final StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + responseObserverRef.set(responseObserver); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + capturedRequests.add(request); + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + } + extProcLatch.countDown(); + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + ExternalProcessor proto = + ExternalProcessor.newBuilder() + .setGrpcService( + GrpcService.newBuilder() + .setGoogleGrpc( + GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///" + uniqueExtProcServerName) + .addChannelCredentialsPlugin( + Any.newBuilder() + .setTypeUrl( + "type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build()) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SEND) + .build()) + .build(); + + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName) + .directExecutor() + .build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + ServerCallHandler nextHandler = + (call, headers) -> { + return new ServerCall.Listener() {}; + }; + + SimpleServerCall dummyCall = new SimpleServerCall(METHOD_SAY_HELLO_RAW); + + try { + @SuppressWarnings("unchecked") + ServerCall.Listener serverListener = + (ServerCall.Listener) + (ServerCall.Listener) + interceptor.interceptCall(dummyCall, new Metadata(), nextHandler); + + // Invoke onMessage() while the call is IDLE (headers response has not been sent) + byte[] messageBytes = "hello".getBytes(StandardCharsets.UTF_8); + serverListener.onMessage(new ByteArrayInputStream(messageBytes)); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + assertThat(capturedRequests).hasSize(2); + + // First request (RequestHeaders) should have protocol_config + ProcessingRequest firstReq = capturedRequests.get(0); + assertThat(firstReq.hasRequestHeaders()).isTrue(); + assertThat(firstReq.hasProtocolConfig()).isTrue(); + assertThat(firstReq.getProtocolConfig().getRequestBodyMode()) + .isEqualTo(ProcessingMode.BodySendMode.GRPC); + assertThat(firstReq.getProtocolConfig().getResponseBodyMode()) + .isEqualTo(ProcessingMode.BodySendMode.GRPC); + + // Second request (RequestBody) should NOT have protocol_config + ProcessingRequest secondReq = capturedRequests.get(1); + assertThat(secondReq.hasRequestBody()).isTrue(); + assertThat(secondReq.hasProtocolConfig()).isFalse(); + } finally { + StreamObserver responseObserver = responseObserverRef.get(); + if (responseObserver != null) { + responseObserver.onCompleted(); + } + channelManager.close(); + } + } + + @Test + public void protocolConfig_onBody() throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + + final CountDownLatch extProcLatch = new CountDownLatch(1); + final List capturedRequests = + Collections.synchronizedList(new ArrayList<>()); + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + final StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + responseObserverRef.set(responseObserver); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + capturedRequests.add(request); + if (request.hasRequestBody()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestBody(BodyResponse.newBuilder().build()) + .build()); + } + extProcLatch.countDown(); + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + ExternalProcessor proto = + createBaseProto(uniqueExtProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setRequestBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseBodyMode(ProcessingMode.BodySendMode.NONE) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SKIP) + .build()) + .build(); + + ExternalProcessorFilterConfig filterConfig = + provider.parseFilterConfig(Any.pack(proto), filterContext).config; + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName) + .directExecutor() + .build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + ServerCallHandler nextHandler = + (call, headers) -> { + return new ServerCall.Listener() {}; + }; + + SimpleServerCall dummyCall = new SimpleServerCall(METHOD_SAY_HELLO_RAW); + try { + ServerCall.Listener serverListener = + interceptor.interceptCall(dummyCall, new Metadata(), nextHandler); + + byte[] messageBytes = "hello".getBytes(StandardCharsets.UTF_8); + serverListener.onMessage(new ByteArrayInputStream(messageBytes)); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(capturedRequests).hasSize(1); + + ProcessingRequest firstReq = capturedRequests.get(0); + assertThat(firstReq.hasRequestBody()).isTrue(); + assertThat(firstReq.hasProtocolConfig()).isTrue(); + assertThat(firstReq.getProtocolConfig().getRequestBodyMode()) + .isEqualTo(ProcessingMode.BodySendMode.GRPC); + } finally { + StreamObserver responseObserver = responseObserverRef.get(); + if (responseObserver != null) { + responseObserver.onCompleted(); + } + channelManager.close(); + } + } + + @Test + public void protocolConfig_onResponseHeaders() throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + + final CountDownLatch extProcLatch = new CountDownLatch(1); + final List capturedRequests = + Collections.synchronizedList(new ArrayList<>()); + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + final StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + responseObserverRef.set(responseObserver); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + capturedRequests.add(request); + if (request.hasResponseHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders(HeadersResponse.newBuilder().build()) + .build()); + } + extProcLatch.countDown(); + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + ExternalProcessor proto = + createBaseProto(uniqueExtProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setRequestBodyMode(ProcessingMode.BodySendMode.NONE) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SEND) + .setResponseBodyMode(ProcessingMode.BodySendMode.NONE) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SKIP) + .build()) + .build(); + + ExternalProcessorFilterConfig filterConfig = + provider.parseFilterConfig(Any.pack(proto), filterContext).config; + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName) + .directExecutor() + .build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + assertThat(capturedRequests).hasSize(1); + ProcessingRequest firstReq = capturedRequests.get(0); + assertThat(firstReq.hasResponseHeaders()).isTrue(); + assertThat(firstReq.hasProtocolConfig()).isTrue(); + assertThat(firstReq.getProtocolConfig().getRequestBodyMode()) + .isEqualTo(ProcessingMode.BodySendMode.NONE); + assertThat(firstReq.getProtocolConfig().getResponseBodyMode()) + .isEqualTo(ProcessingMode.BodySendMode.NONE); + } + + @Test + public void protocolConfig_onResponseBody() throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + + final CountDownLatch extProcLatch = new CountDownLatch(1); + final List capturedRequests = + Collections.synchronizedList(new ArrayList<>()); + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + final StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + responseObserverRef.set(responseObserver); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + capturedRequests.add(request); + if (request.hasResponseBody()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseBody( + BodyResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setBody( + request.getResponseBody().getBody()) + .build()) + .build()) + .build()) + .build()) + .build()); + } else if (request.hasResponseTrailers()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseTrailers(TrailersResponse.newBuilder().build()) + .build()); + } + extProcLatch.countDown(); + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + ExternalProcessor proto = + createBaseProto(uniqueExtProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setRequestBodyMode(ProcessingMode.BodySendMode.NONE) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SEND) + .build()) + .build(); + + ExternalProcessorFilterConfig filterConfig = + provider.parseFilterConfig(Any.pack(proto), filterContext).config; + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName) + .directExecutor() + .build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + assertThat(capturedRequests.size()).isAtLeast(1); + ProcessingRequest firstReq = capturedRequests.get(0); + assertThat(firstReq.hasResponseBody()).isTrue(); + assertThat(firstReq.hasProtocolConfig()).isTrue(); + assertThat(firstReq.getProtocolConfig().getRequestBodyMode()) + .isEqualTo(ProcessingMode.BodySendMode.NONE); + assertThat(firstReq.getProtocolConfig().getResponseBodyMode()) + .isEqualTo(ProcessingMode.BodySendMode.GRPC); + + for (int i = 1; i < capturedRequests.size(); i++) { + assertThat(capturedRequests.get(i).hasProtocolConfig()).isFalse(); + } + } + + @Test + public void protocolConfig_onResponseTrailers() throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + + final CountDownLatch extProcLatch = new CountDownLatch(1); + final List capturedRequests = + Collections.synchronizedList(new ArrayList<>()); + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + final StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + responseObserverRef.set(responseObserver); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + capturedRequests.add(request); + if (request.hasResponseTrailers()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseTrailers(TrailersResponse.newBuilder().build()) + .build()); + } + extProcLatch.countDown(); + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + ExternalProcessor proto = + createBaseProto(uniqueExtProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setRequestBodyMode(ProcessingMode.BodySendMode.NONE) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseBodyMode(ProcessingMode.BodySendMode.NONE) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SEND) + .build()) + .build(); + + ExternalProcessorFilterConfig filterConfig = + provider.parseFilterConfig(Any.pack(proto), filterContext).config; + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName) + .directExecutor() + .build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + assertThat(capturedRequests).hasSize(1); + ProcessingRequest firstReq = capturedRequests.get(0); + assertThat(firstReq.hasResponseTrailers()).isTrue(); + assertThat(firstReq.hasProtocolConfig()).isTrue(); + assertThat(firstReq.getProtocolConfig().getRequestBodyMode()) + .isEqualTo(ProcessingMode.BodySendMode.NONE); + assertThat(firstReq.getProtocolConfig().getResponseBodyMode()) + .isEqualTo(ProcessingMode.BodySendMode.NONE); + } + + // ============================================================================ + // Category 4: GrpcService Initial Metadata + // ============================================================================ + + @Test + public void givenGrpcServiceWithInitialMetadata_whenCallIntercepted_thenSendsMetadata() + throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + + final AtomicReference capturedHeaders = new AtomicReference<>(); + final CountDownLatch extProcStartedLatch = new CountDownLatch(1); + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + final StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + + ServerInterceptor headerCapturingInterceptor = + new ServerInterceptor() { + @Override + public ServerCall.Listener interceptCall( + ServerCall call, Metadata headers, ServerCallHandler next) { + capturedHeaders.set(headers); + extProcStartedLatch.countDown(); + return next.startCall(call, headers); + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(ServerInterceptors.intercept(extProcImpl, headerCapturingInterceptor)) + .directExecutor() + .build() + .start()); + + // Config with initial metadata + ExternalProcessor.Builder protoBuilder = createBaseProto(uniqueExtProcServerName); + protoBuilder.setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SEND) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .build()); + protoBuilder + .getGrpcServiceBuilder() + .addInitialMetadata( + io.envoyproxy.envoy.config.core.v3.HeaderValue.newBuilder() + .setKey("x-init-key") + .setValue("init-val") + .build()) + .addInitialMetadata( + io.envoyproxy.envoy.config.core.v3.HeaderValue.newBuilder() + .setKey("x-bin-key-bin") + .setRawValue(ByteString.copyFrom(new byte[] {1, 2, 3})) + .build()); + ExternalProcessor proto = protoBuilder.build(); + + ExternalProcessorFilterConfig filterConfig = + provider.parseFilterConfig(Any.pack(proto), filterContext).config; + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName) + .directExecutor() + .build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + final AtomicReference callStatus = new AtomicReference<>(); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callStatus.set(status); + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + boolean sidecarAwaited = extProcStartedLatch.await(5, TimeUnit.SECONDS); + boolean completedAwaited = callCompletedLatch.await(5, TimeUnit.SECONDS); + + assertThat(sidecarAwaited).isTrue(); + assertThat(completedAwaited).isTrue(); + assertThat(callStatus.get().isOk()).isTrue(); + + assertThat(capturedHeaders.get()).isNotNull(); + assertThat( + capturedHeaders + .get() + .get(Metadata.Key.of("x-init-key", Metadata.ASCII_STRING_MARSHALLER))) + .isEqualTo("init-val"); + assertThat( + capturedHeaders + .get() + .get(Metadata.Key.of("x-bin-key-bin", Metadata.BINARY_BYTE_MARSHALLER))) + .isEqualTo(new byte[] {1, 2, 3}); + } + + // ============================================================================ + // Category 5: Request attributes propagation + // ============================================================================ + + @Test + public void requestAttributes_onHeaders() throws Exception { + final AtomicReference capturedRequest = new AtomicReference<>(); + final CountDownLatch sidecarLatch = new CountDownLatch(1); + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + capturedRequest.set(request); + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + sidecarLatch.countDown(); + } else if (request.hasResponseHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders(HeadersResponse.newBuilder().build()) + .build()); + } else if (request.hasRequestBody()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestBody( + io.envoyproxy.envoy.service.ext_proc.v3.BodyResponse.newBuilder() + .build()) + .build()); + } else if (request.hasResponseBody()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseBody( + io.envoyproxy.envoy.service.ext_proc.v3.BodyResponse.newBuilder() + .build()) + .build()); + } else if (request.hasRequestTrailers()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestTrailers( + io.envoyproxy.envoy.service.ext_proc.v3.TrailersResponse.newBuilder() + .build()) + .build()); + } else if (request.hasResponseTrailers()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseTrailers( + io.envoyproxy.envoy.service.ext_proc.v3.TrailersResponse.newBuilder() + .build()) + .build()); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + // Config with request attributes requested + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SEND) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SKIP) + .build()) + .addRequestAttributes("request.path") + .addRequestAttributes("request.host") + .addRequestAttributes("request.method") + .addRequestAttributes("request.query") + .build(); + + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(sidecarLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + ProcessingRequest request = capturedRequest.get(); + java.util.Map attributes = request.getAttributesMap(); + assertThat(attributes.get("request.path").getFieldsOrThrow("").getStringValue()) + .isEqualTo("/test.TestService/SayHello"); + assertThat(attributes.get("request.host").getFieldsOrThrow("").getStringValue()) + .isEqualTo(dataPlaneChannel.authority()); + assertThat(attributes.get("request.method").getFieldsOrThrow("").getStringValue()) + .isEqualTo("POST"); + assertThat(attributes.get("request.query").getFieldsOrThrow("").getStringValue()).isEqualTo(""); + } + + @Test + public void requestAttributes_metadata() throws Exception { + final AtomicReference capturedRequest = new AtomicReference<>(); + final CountDownLatch sidecarLatch = new CountDownLatch(1); + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + capturedRequest.set(request); + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + sidecarLatch.countDown(); + } else if (request.hasResponseHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders(HeadersResponse.newBuilder().build()) + .build()); + } else if (request.hasRequestBody()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestBody( + io.envoyproxy.envoy.service.ext_proc.v3.BodyResponse.newBuilder() + .build()) + .build()); + } else if (request.hasResponseBody()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseBody( + io.envoyproxy.envoy.service.ext_proc.v3.BodyResponse.newBuilder() + .build()) + .build()); + } else if (request.hasRequestTrailers()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestTrailers( + io.envoyproxy.envoy.service.ext_proc.v3.TrailersResponse.newBuilder() + .build()) + .build()); + } else if (request.hasResponseTrailers()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseTrailers( + io.envoyproxy.envoy.service.ext_proc.v3.TrailersResponse.newBuilder() + .build()) + .build()); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + // Config with metadata attributes requested + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SEND) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SKIP) + .build()) + .addRequestAttributes("request.referer") + .addRequestAttributes("request.useragent") + .addRequestAttributes("request.id") + .addRequestAttributes("request.headers") + .build(); + + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + Metadata headers = new Metadata(); + headers.put(Metadata.Key.of("referer", Metadata.ASCII_STRING_MARSHALLER), "http://google.com"); + headers.put(Metadata.Key.of("user-agent", Metadata.ASCII_STRING_MARSHALLER), "custom-ua"); + headers.put(Metadata.Key.of("x-request-id", Metadata.ASCII_STRING_MARSHALLER), "req-123"); + headers.put(Metadata.Key.of("custom-header", Metadata.ASCII_STRING_MARSHALLER), "val"); + headers.put( + Metadata.Key.of("x-bin-key-bin", Metadata.BINARY_BYTE_MARSHALLER), new byte[] {1, 2}); + + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + headers); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(sidecarLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + ProcessingRequest request = capturedRequest.get(); + java.util.Map attributes = request.getAttributesMap(); + assertThat(attributes.get("request.referer").getFieldsOrThrow("").getStringValue()) + .isEqualTo("http://google.com"); + assertThat(attributes.get("request.useragent").getFieldsOrThrow("").getStringValue()) + .contains("custom-ua"); + assertThat(attributes.get("request.id").getFieldsOrThrow("").getStringValue()) + .isEqualTo("req-123"); + + com.google.protobuf.Struct headersStruct = attributes.get("request.headers"); + assertThat(headersStruct.getFieldsOrThrow("custom-header").getStringValue()).isEqualTo("val"); + assertThat(headersStruct.getFieldsOrThrow("x-bin-key-bin").getStringValue()).isEqualTo("AQI"); + } + + @Test + public void requestAttributes_onBody() throws Exception { + final java.util.List capturedRequests = + java.util.Collections.synchronizedList(new java.util.ArrayList<>()); + final CountDownLatch sidecarLatch = new CountDownLatch(1); + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + capturedRequests.add(request); + if (request.hasRequestBody()) { + BodyResponse.Builder bodyResponse = BodyResponse.newBuilder(); + if (request.getRequestBody().getBody().isEmpty() + && (request.getRequestBody().getEndOfStream() + || request.getRequestBody().getEndOfStreamWithoutMessage())) { + bodyResponse.setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setEndOfStream( + request.getRequestBody().getEndOfStream()) + .setEndOfStreamWithoutMessage( + request + .getRequestBody() + .getEndOfStreamWithoutMessage()) + .build()) + .build()) + .build()); + } else { + bodyResponse.setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setBody(request.getRequestBody().getBody()) + .setEndOfStream( + request.getRequestBody().getEndOfStream()) + .build()) + .build()) + .build()); + } + responseObserver.onNext( + ProcessingResponse.newBuilder().setRequestBody(bodyResponse).build()); + sidecarLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setRequestBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SKIP) + .build()) + .addRequestAttributes("request.path") + .addRequestAttributes("request.host") + .build(); + + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(sidecarLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + assertThat(capturedRequests.size()).isAtLeast(1); + ProcessingRequest firstReq = capturedRequests.get(0); + assertThat(firstReq.hasRequestBody()).isTrue(); + java.util.Map attributes = firstReq.getAttributesMap(); + assertThat(attributes.get("request.path").getFieldsOrThrow("").getStringValue()) + .isEqualTo("/test.TestService/SayHello"); + assertThat(attributes.get("request.host").getFieldsOrThrow("").getStringValue()) + .isEqualTo(dataPlaneChannel.authority()); + + for (int i = 1; i < capturedRequests.size(); i++) { + assertThat(capturedRequests.get(i).getAttributesCount()).isEqualTo(0); + } + } + + @Test + public void requestAttributes_notSent() throws Exception { + final AtomicReference capturedRequest = new AtomicReference<>(); + final CountDownLatch sidecarLatch = new CountDownLatch(1); + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasResponseHeaders()) { + capturedRequest.set(request); + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders(HeadersResponse.newBuilder().build()) + .build()); + sidecarLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setRequestBodyMode(ProcessingMode.BodySendMode.NONE) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SEND) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SKIP) + .build()) + .addRequestAttributes("request.path") + .addRequestAttributes("request.host") + .build(); + + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(sidecarLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + ProcessingRequest request = capturedRequest.get(); + assertThat(request.hasResponseHeaders()).isTrue(); + assertThat(request.getAttributesCount()).isEqualTo(0); + } + + // ============================================================================ + // Category 6: Request Header Processing + // ============================================================================ + @Test + public void givenRequestHeaderModeSend_whenStartCallCalled_thenCallIsBuffered() throws Exception { + // Configure GRPC request body mode + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestBodyMode(ProcessingMode.BodySendMode.GRPC) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final AtomicReference receivedRequestRef = new AtomicReference<>(); + final CountDownLatch requestLatch = new CountDownLatch(2); // Expect 1 for headers, 1 for body + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + responseObserverRef.set(responseObserver); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestBody()) { + receivedRequestRef.set(request); + } + requestLatch.countDown(); + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + SimpleServerCall dummyCall = new SimpleServerCall(METHOD_SAY_HELLO_RAW); + + final AtomicBoolean startCallCalled = new AtomicBoolean(); + ServerCallHandler dummyNext = + new ServerCallHandler() { + @Override + public ServerCall.Listener startCall( + ServerCall call, Metadata headers) { + startCallCalled.set(true); + return new ServerCall.Listener() {}; + } + }; + + @SuppressWarnings("unchecked") + ServerCall.Listener serverListener = + (ServerCall.Listener) + (ServerCall.Listener) + interceptor.interceptCall(dummyCall, new Metadata(), dummyNext); + + // Invocate onMessage() while the call is IDLE (headers response has not been sent) + byte[] messageBytes = "hello".getBytes(StandardCharsets.UTF_8); + serverListener.onMessage(new ByteArrayInputStream(messageBytes)); + + // Wait and verify that the external processor immediately receives both the headers request and + // the body request + boolean receivedInTime = requestLatch.await(5, TimeUnit.SECONDS); + assertThat(receivedInTime).isTrue(); + + ProcessingRequest receivedRequest = receivedRequestRef.get(); + assertThat(receivedRequest).isNotNull(); + assertThat(receivedRequest.hasRequestBody()).isTrue(); + assertThat(receivedRequest.getRequestBody().getBody().toStringUtf8()).isEqualTo("hello"); + + // Assert that the data plane call was not started yet (buffered) + assertThat(startCallCalled.get()).isFalse(); + + // Clean up control stream to allow resources to be released cleanly + StreamObserver responseObserver = responseObserverRef.get(); + if (responseObserver != null) { + responseObserver.onCompleted(); + } + } + + @Test + public void givenRequestHeaderModeSend_whenExtProcRespondsWithMutations_thenCallIsActivated() + throws Exception { + // Configure GRPC request body mode + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestBodyMode(ProcessingMode.BodySendMode.GRPC) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + final CountDownLatch requestLatch = + new CountDownLatch(1); // Just wait for headers request to be processed + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + responseObserverRef.set(responseObserver); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + requestLatch.countDown(); + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + SimpleServerCall dummyCall = new SimpleServerCall(METHOD_SAY_HELLO_RAW); + + final AtomicBoolean startCallCalled = new AtomicBoolean(); + ServerCallHandler dummyNext = + new ServerCallHandler() { + @Override + public ServerCall.Listener startCall( + ServerCall call, Metadata headers) { + startCallCalled.set(true); + return new ServerCall.Listener() {}; + } + }; + + @SuppressWarnings("unchecked") + ServerCall.Listener serverListener = + (ServerCall.Listener) + (ServerCall.Listener) + interceptor.interceptCall(dummyCall, new Metadata(), dummyNext); + + // Invocate onMessage() while the call is IDLE (headers response has not been sent) + byte[] messageBytes = "hello".getBytes(StandardCharsets.UTF_8); + serverListener.onMessage(new ByteArrayInputStream(messageBytes)); + + boolean receivedInTime = requestLatch.await(5, TimeUnit.SECONDS); + assertThat(receivedInTime).isTrue(); + + // Assert that the data plane call was not started yet + assertThat(startCallCalled.get()).isFalse(); + + // Clean up control stream by completing it, which triggers fallback activation on the data + // plane call + StreamObserver responseObserver = responseObserverRef.get(); + assertThat(responseObserver).isNotNull(); + responseObserver.onNext( + ProcessingResponse.newBuilder().setRequestDrain(true).build()); + responseObserver.onCompleted(); + + // Verify that the call is now activated + assertThat(startCallCalled.get()).isTrue(); + } + + @Test + public void serverInterceptor_headerMutation_addsHeader() throws Exception { + ExternalProcessor proto = createBaseProto(extProcServerName).build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = new CountDownLatch(1); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders( + HeadersResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setHeaderMutation( + HeaderMutation.newBuilder() + .addSetHeaders( + io.envoyproxy.envoy.config.core.v3 + .HeaderValueOption.newBuilder() + .setHeader( + io.envoyproxy.envoy.config.core.v3 + .HeaderValue.newBuilder() + .setKey("x-mutated-header") + .setValue("mutated-value") + .build()) + .build()) + .build()) + .build()) + .build()) + .build()); + responseObserver.onCompleted(); + extProcLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicReference receivedHeaders = new AtomicReference<>(); + ServerInterceptor capturingInterceptor = + new ServerInterceptor() { + @Override + public ServerCall.Listener interceptCall( + ServerCall call, Metadata headers, ServerCallHandler next) { + receivedHeaders.set(headers); + return next.startCall(call, headers); + } + }; + + startDataPlane(interceptor, capturingInterceptor); + + Metadata initialHeaders = new Metadata(); + initialHeaders.put( + Metadata.Key.of("x-initial-header", Metadata.ASCII_STRING_MARSHALLER), "initial-value"); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + initialHeaders); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(receivedHeaders.get()).isNotNull(); + assertThat( + receivedHeaders + .get() + .get(Metadata.Key.of("x-mutated-header", Metadata.ASCII_STRING_MARSHALLER))) + .isEqualTo("mutated-value"); + } + + @Test + public void serverInterceptor_requestHeaderModeSkip_doesNotSendHeaders() throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SKIP) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = new CountDownLatch(1); + final AtomicBoolean requestHeadersReceived = new AtomicBoolean(false); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + requestHeadersReceived.set(true); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + extProcLatch.countDown(); + } + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicReference receivedHeaders = new AtomicReference<>(); + ServerInterceptor capturingInterceptor = + new ServerInterceptor() { + @Override + public ServerCall.Listener interceptCall( + ServerCall call, Metadata headers, ServerCallHandler next) { + receivedHeaders.set(headers); + return next.startCall(call, headers); + } + }; + + startDataPlane(interceptor, capturingInterceptor); + + Metadata initialHeaders = new Metadata(); + initialHeaders.put( + Metadata.Key.of("x-initial-header", Metadata.ASCII_STRING_MARSHALLER), "initial-value"); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + initialHeaders); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(requestHeadersReceived.get()).isFalse(); + assertThat(receivedHeaders.get()).isNotNull(); + assertThat( + receivedHeaders + .get() + .get(Metadata.Key.of("x-initial-header", Metadata.ASCII_STRING_MARSHALLER))) + .isEqualTo("initial-value"); + } + + // ============================================================================ + // Category 7: Body Mutation: Inbound/Request (GRPC Mode) + // ============================================================================ + + @Test + public void givenRequestBodyModeGrpc_whenMessageReceived_thenMessageSentToExtProc() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setRequestBodyMode(ProcessingMode.BodySendMode.GRPC) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch bodySentLatch = new CountDownLatch(1); + final AtomicReference capturedRequest = new AtomicReference<>(); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestBody()) { + if (capturedRequest.get() == null + && !request.getRequestBody().getBody().isEmpty()) { + capturedRequest.set(request); + bodySentLatch.countDown(); + } + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestBody( + BodyResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setEndOfStream( + request + .getRequestBody() + .getEndOfStream()) + .build()) + .build()) + .build()) + .build()) + .build()); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + clientCall.start(new io.grpc.ClientCall.Listener() {}, new Metadata()); + clientCall.request(1); + clientCall.sendMessage( + new ByteArrayInputStream("Hello World".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(bodySentLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(capturedRequest.get().getRequestBody().getBody().toStringUtf8()) + .contains("Hello World"); + + clientCall.cancel("Cleanup", null); + channelManager.close(); + } + + @Test + public void givenRequestBodyModeGrpc_whenExtProcRespondsWithMutatedBody_thenMutatedBodyForwarded() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setRequestBodyMode(ProcessingMode.BodySendMode.GRPC) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestBody()) { + BodyResponse.Builder bodyResponse = BodyResponse.newBuilder(); + if (request.getRequestBody().getBody().isEmpty() + && (request.getRequestBody().getEndOfStream() + || request.getRequestBody().getEndOfStreamWithoutMessage())) { + bodyResponse.setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setEndOfStream( + request.getRequestBody().getEndOfStream()) + .setEndOfStreamWithoutMessage( + request + .getRequestBody() + .getEndOfStreamWithoutMessage()) + .build()) + .build()) + .build()); + } else { + bodyResponse.setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setBody( + ByteString.copyFromUtf8("Mutated Request Body")) + .setEndOfStream( + request.getRequestBody().getEndOfStream()) + .build()) + .build()) + .build()); + } + responseObserver.onNext( + ProcessingResponse.newBuilder().setRequestBody(bodyResponse).build()); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicReference receivedBody = new AtomicReference<>(); + final CountDownLatch dataPlaneLatch = new CountDownLatch(1); + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + try { + byte[] bytes = ByteStreams.toByteArray(request); + receivedBody.set(new String(bytes, StandardCharsets.UTF_8)); + } catch (Exception e) { + receivedBody.set("Error reading: " + e.getMessage()); + } + responseObserver.onNext( + new ByteArrayInputStream("Hello Back".getBytes(StandardCharsets.UTF_8))); + responseObserver.onCompleted(); + dataPlaneLatch.countDown(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + clientCall.start(new io.grpc.ClientCall.Listener() {}, new Metadata()); + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("Original".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(dataPlaneLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(receivedBody.get()).isEqualTo("Mutated Request Body"); + + clientCall.cancel("Cleanup", null); + channelManager.close(); + } + + @Test + public void + givenRequestBodyModeGrpc_whenExtProcRespondsWithEmptyBody_thenEmptyMessageIsDelivered() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setRequestBodyMode(ProcessingMode.BodySendMode.GRPC) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestBody()) { + BodyResponse.Builder bodyResponse = BodyResponse.newBuilder(); + if (request.getRequestBody().getBody().isEmpty() + && (request.getRequestBody().getEndOfStream() + || request.getRequestBody().getEndOfStreamWithoutMessage())) { + bodyResponse.setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setEndOfStream( + request.getRequestBody().getEndOfStream()) + .setEndOfStreamWithoutMessage( + request + .getRequestBody() + .getEndOfStreamWithoutMessage()) + .build()) + .build()) + .build()); + } else { + bodyResponse.setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setBody(ByteString.EMPTY) // Mutate to EMPTY! + .build()) + .build()) + .build()); + } + responseObserver.onNext( + ProcessingResponse.newBuilder().setRequestBody(bodyResponse).build()); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicReference receivedBody = new AtomicReference<>(); + final CountDownLatch dataPlaneLatch = new CountDownLatch(1); + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + try { + byte[] bytes = ByteStreams.toByteArray(request); + receivedBody.set(new String(bytes, StandardCharsets.UTF_8)); + } catch (Exception e) { + receivedBody.set("Error reading: " + e.getMessage()); + } + responseObserver.onNext( + new ByteArrayInputStream("Hello Back".getBytes(StandardCharsets.UTF_8))); + responseObserver.onCompleted(); + dataPlaneLatch.countDown(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + clientCall.start(new io.grpc.ClientCall.Listener() {}, new Metadata()); + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("Original".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(dataPlaneLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(receivedBody.get()).isEmpty(); // Assert it is EMPTY! + + clientCall.cancel("Cleanup", null); + channelManager.close(); + } + + @Test + public void givenExtProcSignaledEndOfStream_whenMoreMessagesReceived_thenMessagesDiscarded() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setRequestBodyMode(ProcessingMode.BodySendMode.GRPC) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final AtomicInteger sidecarMessages = new AtomicInteger(0); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestBody()) { + sidecarMessages.incrementAndGet(); + boolean triggerEos = + request.getRequestBody().getBody().toStringUtf8().equals("Trigger EOS"); + BodyResponse.Builder bodyResponse = BodyResponse.newBuilder(); + if (triggerEos + || (request.getRequestBody().getBody().isEmpty() + && (request.getRequestBody().getEndOfStream() + || request.getRequestBody().getEndOfStreamWithoutMessage()))) { + bodyResponse.setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setBody(request.getRequestBody().getBody()) + .setEndOfStream( + triggerEos + || request.getRequestBody().getEndOfStream()) + .setEndOfStreamWithoutMessage( + request + .getRequestBody() + .getEndOfStreamWithoutMessage()) + .build()) + .build()) + .build()); + } else { + bodyResponse.setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setEndOfStream( + request.getRequestBody().getEndOfStream()) + .build()) + .build()) + .build()); + } + responseObserver.onNext( + ProcessingResponse.newBuilder().setRequestBody(bodyResponse).build()); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicInteger dataPlaneMessages = new AtomicInteger(0); + final CountDownLatch dataPlaneHalfCloseLatch = new CountDownLatch(1); + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public StreamObserver sayHelloClientStreaming( + final StreamObserver responseObserver) { + return new StreamObserver() { + @Override + public void onNext(InputStream request) { + try { + byte[] unused = ByteStreams.toByteArray(request); + dataPlaneMessages.incrementAndGet(); + } catch (Exception e) { + // Ignore + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onNext( + new ByteArrayInputStream("Hello Back".getBytes(StandardCharsets.UTF_8))); + responseObserver.onCompleted(); + dataPlaneHalfCloseLatch.countDown(); + } + }; + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_CLIENT_STREAMING, io.grpc.CallOptions.DEFAULT); + + clientCall.start(new io.grpc.ClientCall.Listener() {}, new Metadata()); + clientCall.request(10); + clientCall.sendMessage( + new ByteArrayInputStream("Trigger EOS".getBytes(StandardCharsets.UTF_8))); + + // We need to wait until the sidecar processes the message and signals endOfStream. + // When endOfStream is received by data plane, it calls delegate.onHalfClose() + assertThat(dataPlaneHalfCloseLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(dataPlaneMessages.get()).isEqualTo(1); + assertThat(sidecarMessages.get()).isEqualTo(1); + + // Now send another message. It should be discarded! + clientCall.sendMessage(new ByteArrayInputStream("Too late".getBytes(StandardCharsets.UTF_8))); + + // Wait a little bit to make sure it is processed and discarded + Thread.sleep(100); + + assertThat(dataPlaneMessages.get()).isEqualTo(1); + assertThat(sidecarMessages.get()).isEqualTo(1); + + clientCall.cancel("Cleanup", null); + channelManager.close(); + } + + @Test + public void givenRequestBodyModeNone_whenSendMessageCalled_thenMessageSentDirectlyToDataPlane() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setRequestBodyMode(ProcessingMode.BodySendMode.NONE) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final AtomicInteger extProcBodyCount = new AtomicInteger(0); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestBody()) { + extProcBodyCount.incrementAndGet(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicReference receivedBody = new AtomicReference<>(); + final CountDownLatch dataPlaneLatch = new CountDownLatch(1); + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + try { + byte[] bytes = ByteStreams.toByteArray(request); + receivedBody.set(new String(bytes, StandardCharsets.UTF_8)); + } catch (Exception e) { + receivedBody.set("Error reading: " + e.getMessage()); + } + responseObserver.onNext( + new ByteArrayInputStream("Hello Back".getBytes(StandardCharsets.UTF_8))); + responseObserver.onCompleted(); + dataPlaneLatch.countDown(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + clientCall.start(new io.grpc.ClientCall.Listener() {}, new Metadata()); + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("Original".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(dataPlaneLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(receivedBody.get()).isEqualTo("Original"); + assertThat(extProcBodyCount.get()).isEqualTo(0); + + clientCall.cancel("Cleanup", null); + channelManager.close(); + } + + // ============================================================================ + // Category 8: Response Header Mutation + // ============================================================================ + + @Test + public void serverInterceptor_responseHeaderMutation_mutatesHeader() throws Exception { + ExternalProcessor proto = createBaseProto(extProcServerName).build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = + new CountDownLatch(2); // 1 for request headers, 1 for response headers + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders( + HeadersResponse.newBuilder() + .setResponse(CommonResponse.newBuilder().build()) + .build()) + .build()); + extProcLatch.countDown(); + } else if (request.hasResponseHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders( + HeadersResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setHeaderMutation( + HeaderMutation.newBuilder() + .addSetHeaders( + io.envoyproxy.envoy.config.core.v3 + .HeaderValueOption.newBuilder() + .setHeader( + io.envoyproxy.envoy.config.core.v3 + .HeaderValue.newBuilder() + .setKey( + "x-mutated-response-header") + .setValue( + "mutated-response-value") + .build()) + .build()) + .build()) + .build()) + .build()) + .build()); + responseObserver.onCompleted(); + extProcLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + ServerInterceptor headersInterceptor = + new ServerInterceptor() { + @Override + public ServerCall.Listener interceptCall( + ServerCall call, Metadata headers, ServerCallHandler next) { + return next.startCall( + new ForwardingServerCall.SimpleForwardingServerCall(call) { + @Override + public void sendHeaders(Metadata responseHeaders) { + responseHeaders.put( + Metadata.Key.of( + "x-initial-response-header", Metadata.ASCII_STRING_MARSHALLER), + "initial-response-value"); + super.sendHeaders(responseHeaders); + } + }, + headers); + } + }; + + startDataPlane(interceptor, headersInterceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final AtomicReference receivedResponseHeaders = new AtomicReference<>(); + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onHeaders(Metadata headers) { + receivedResponseHeaders.set(headers); + } + + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(receivedResponseHeaders.get()).isNotNull(); + assertThat( + receivedResponseHeaders + .get() + .get( + Metadata.Key.of("x-mutated-response-header", Metadata.ASCII_STRING_MARSHALLER))) + .isEqualTo("mutated-response-value"); + assertThat( + receivedResponseHeaders + .get() + .get( + Metadata.Key.of("x-initial-response-header", Metadata.ASCII_STRING_MARSHALLER))) + .isEqualTo("initial-response-value"); + } + + @Test + public void givenResponseHeaderModeSkip_responseHeadersSentDirectlyDownstream() throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SEND) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + Metadata.Key customKey = + Metadata.Key.of("custom-response-header", Metadata.ASCII_STRING_MARSHALLER); + + final java.util.concurrent.atomic.AtomicBoolean responseHeadersReceived = + new java.util.concurrent.atomic.AtomicBoolean(false); + final CountDownLatch requestHeadersLatch = new CountDownLatch(1); + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + ProcessingResponse.Builder response = ProcessingResponse.newBuilder(); + if (request.hasRequestHeaders()) { + response.setRequestHeaders( + HeadersResponse.newBuilder() + .setResponse(CommonResponse.newBuilder().build()) + .build()); + requestHeadersLatch.countDown(); + } else if (request.hasResponseHeaders()) { + responseHeadersReceived.set(true); + } + responseObserver.onNext(response.build()); + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + ServerInterceptor headersInterceptor = + new ServerInterceptor() { + @Override + public ServerCall.Listener interceptCall( + ServerCall call, Metadata headers, ServerCallHandler next) { + return next.startCall( + new ForwardingServerCall.SimpleForwardingServerCall(call) { + @Override + public void sendHeaders(Metadata responseHeaders) { + responseHeaders.put(customKey, "custom-value"); + super.sendHeaders(responseHeaders); + } + }, + headers); + } + }; + + startDataPlane(interceptor, headersInterceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final AtomicReference receivedResponseHeaders = new AtomicReference<>(); + final CountDownLatch headersLatch = new CountDownLatch(1); + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onHeaders(Metadata headers) { + receivedResponseHeaders.set(headers); + headersLatch.countDown(); + } + + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(requestHeadersLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(headersLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + assertThat(receivedResponseHeaders.get()).isNotNull(); + assertThat(receivedResponseHeaders.get().get(customKey)).isEqualTo("custom-value"); + + Thread.sleep(500); + assertThat(responseHeadersReceived.get()).isFalse(); + + channelManager.close(); + } + + // ============================================================================ + // Category 9: Body Mutation: Outbound/Response (GRPC Mode) + // ============================================================================ + + @Test + public void serverInterceptor_responseHeaderModeSkip_doesNotSendResponseHeaders() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = new CountDownLatch(1); // 1 for request headers + final AtomicBoolean responseHeadersReceived = new AtomicBoolean(false); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders( + HeadersResponse.newBuilder() + .setResponse(CommonResponse.newBuilder().build()) + .build()) + .build()); + extProcLatch.countDown(); + } else if (request.hasResponseHeaders()) { + responseHeadersReceived.set(true); + } else if (request.hasResponseBody()) { + responseObserver.onCompleted(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + ServerInterceptor headersInterceptor = + new ServerInterceptor() { + @Override + public ServerCall.Listener interceptCall( + ServerCall call, Metadata headers, ServerCallHandler next) { + return next.startCall( + new ForwardingServerCall.SimpleForwardingServerCall(call) { + @Override + public void sendHeaders(Metadata responseHeaders) { + responseHeaders.put( + Metadata.Key.of( + "x-initial-response-header", Metadata.ASCII_STRING_MARSHALLER), + "initial-response-value"); + super.sendHeaders(responseHeaders); + } + }, + headers); + } + }; + + startDataPlane(interceptor, headersInterceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final AtomicReference receivedResponseHeaders = new AtomicReference<>(); + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onHeaders(Metadata headers) { + receivedResponseHeaders.set(headers); + } + + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(responseHeadersReceived.get()).isFalse(); + assertThat(receivedResponseHeaders.get()).isNotNull(); + assertThat( + receivedResponseHeaders + .get() + .get( + Metadata.Key.of("x-initial-response-header", Metadata.ASCII_STRING_MARSHALLER))) + .isEqualTo("initial-response-value"); + } + + @Test + public void serverInterceptor_responseBodyModeGrpc_whenOnMessageCalled_thenMessageSentToExtProc() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SEND) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch sidecarBodyLatch = new CountDownLatch(1); + final AtomicReference capturedRequest = new AtomicReference<>(); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasResponseBody()) { + if (capturedRequest.get() == null + && !request.getResponseBody().getBody().isEmpty()) { + capturedRequest.set(request); + sidecarBodyLatch.countDown(); + } + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseBody( + BodyResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setBody( + request.getResponseBody().getBody()) + .build()) + .build()) + .build()) + .build()) + .build()); + } else if (request.hasResponseTrailers()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseTrailers(TrailersResponse.newBuilder().build()) + .build()); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext( + new ByteArrayInputStream("Server Message".getBytes(StandardCharsets.UTF_8))); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch appMessageLatch = new CountDownLatch(1); + final AtomicReference receivedMessage = new AtomicReference<>(); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onMessage(InputStream message) { + try { + receivedMessage.set( + new String(ByteStreams.toByteArray(message), StandardCharsets.UTF_8)); + } catch (Exception e) { + receivedMessage.set("Error: " + e.getMessage()); + } + appMessageLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("Hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(sidecarBodyLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat( + new String( + capturedRequest.get().getResponseBody().getBody().toByteArray(), + StandardCharsets.UTF_8)) + .isEqualTo("Server Message"); + assertThat(appMessageLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(receivedMessage.get()).isEqualTo("Server Message"); + } + + @Test + public void serverInterceptor_responseHeadersAndBodyModeGrpc_succeeds() throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SEND) + .setResponseBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SEND) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final List capturedRequests = + java.util.Collections.synchronizedList(new ArrayList<>()); + final CountDownLatch sidecarFinishedLatch = new CountDownLatch(1); + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + capturedRequests.add(request); + if (request.hasResponseHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders(HeadersResponse.newBuilder().build()) + .build()); + } else if (request.hasResponseBody()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseBody( + BodyResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setBody( + request.getResponseBody().getBody()) + .build()) + .build()) + .build()) + .build()) + .build()); + } else if (request.hasResponseTrailers()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseTrailers(TrailersResponse.newBuilder().build()) + .build()); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + sidecarFinishedLatch.countDown(); + } + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext( + new ByteArrayInputStream("Server Message".getBytes(StandardCharsets.UTF_8))); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch clientCompletedLatch = new CountDownLatch(1); + final AtomicReference closedStatus = new AtomicReference<>(); + final AtomicReference receivedMessage = new AtomicReference<>(); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onMessage(InputStream message) { + try { + receivedMessage.set( + new String(ByteStreams.toByteArray(message), StandardCharsets.UTF_8)); + } catch (Exception e) { + receivedMessage.set("Error: " + e.getMessage()); + } + } + + @Override + public void onClose(Status status, Metadata trailers) { + closedStatus.set(status); + clientCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("Hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(clientCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(closedStatus.get().isOk()).isTrue(); + assertThat(receivedMessage.get()).isEqualTo("Server Message"); + + assertThat(sidecarFinishedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + int responseHeadersCount = 0; + int responseBodyCount = 0; + int responseTrailersCount = 0; + for (ProcessingRequest request : capturedRequests) { + if (request.hasResponseHeaders()) { + responseHeadersCount++; + } else if (request.hasResponseBody()) { + responseBodyCount++; + } else if (request.hasResponseTrailers()) { + responseTrailersCount++; + } + } + assertThat(responseHeadersCount).isEqualTo(1); + assertThat(responseBodyCount).isEqualTo(1); + assertThat(responseTrailersCount).isEqualTo(1); + } + + @Test + public void + serverInterceptor_respBodyModeGrpc_whenExtProcRespondsWithMutatedBody_thenMutatedDelivered() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SEND) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasResponseBody()) { + BodyResponse.Builder bodyResponse = BodyResponse.newBuilder(); + if (request.getResponseBody().getBody().isEmpty() + && (request.getResponseBody().getEndOfStream() + || request.getResponseBody().getEndOfStreamWithoutMessage())) { + bodyResponse.setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse(StreamedBodyResponse.newBuilder().build()) + .build()) + .build()); + } else { + bodyResponse.setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setBody( + ByteString.copyFromUtf8("Mutated Server Message")) + .build()) + .build()) + .build()); + } + responseObserver.onNext( + ProcessingResponse.newBuilder().setResponseBody(bodyResponse).build()); + } else if (request.hasResponseTrailers()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseTrailers(TrailersResponse.newBuilder().build()) + .build()); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext( + new ByteArrayInputStream( + "Original Server Message".getBytes(StandardCharsets.UTF_8))); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch appMessageLatch = new CountDownLatch(1); + final AtomicReference receivedMessage = new AtomicReference<>(); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onMessage(InputStream message) { + try { + receivedMessage.set( + new String(ByteStreams.toByteArray(message), StandardCharsets.UTF_8)); + } catch (Exception e) { + receivedMessage.set("Error: " + e.getMessage()); + } + appMessageLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("Hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(appMessageLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(receivedMessage.get()).isEqualTo("Mutated Server Message"); + } + + @Test + public void + serverInterceptor_responseBodyModeGrpc_whenExtProcRespondsEmpty_thenEmptyMsgDelivered() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SEND) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasResponseBody()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseBody( + BodyResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setBody(ByteString.EMPTY) + .setEndOfStream( + request + .getResponseBody() + .getEndOfStream()) + .build()) + .build()) + .build()) + .build()) + .build()); + } else if (request.hasResponseTrailers()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseTrailers(TrailersResponse.newBuilder().build()) + .build()); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext( + new ByteArrayInputStream("Server Message".getBytes(StandardCharsets.UTF_8))); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch appMessageLatch = new CountDownLatch(1); + final AtomicReference receivedMessage = new AtomicReference<>(); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onMessage(InputStream message) { + try { + receivedMessage.set( + new String(ByteStreams.toByteArray(message), StandardCharsets.UTF_8)); + } catch (Exception e) { + receivedMessage.set("Error: " + e.getMessage()); + } + appMessageLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("Hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(appMessageLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(receivedMessage.get()).isEqualTo(""); + } + + @Test + public void + serverInterceptor_respBodyModeNone_whenServerSendsMessage_thenMessageSentDirectlyToClient() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseBodyMode(ProcessingMode.BodySendMode.NONE) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SEND) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final AtomicInteger extProcBodyCount = new AtomicInteger(0); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasResponseBody()) { + extProcBodyCount.incrementAndGet(); + } else if (request.hasResponseTrailers()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseTrailers(TrailersResponse.newBuilder().build()) + .build()); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext( + new ByteArrayInputStream("Server Message".getBytes(StandardCharsets.UTF_8))); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch appMessageLatch = new CountDownLatch(1); + final AtomicReference receivedMessage = new AtomicReference<>(); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onMessage(InputStream message) { + try { + receivedMessage.set( + new String(ByteStreams.toByteArray(message), StandardCharsets.UTF_8)); + } catch (Exception e) { + receivedMessage.set("Error: " + e.getMessage()); + } + appMessageLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("Hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(appMessageLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(receivedMessage.get()).isEqualTo("Server Message"); + assertThat(extProcBodyCount.get()).isEqualTo(0); + } + + // ============================================================================ + // Category 10: Response Trailers + // ============================================================================ + + @Test + public void + givenResponseTrailerModeSend_whenCallCloses_thenResponseTrailersAndStatusPropagatedToClient() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SEND) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = + new CountDownLatch(2); // 1 for request headers, 1 for response trailers + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders( + HeadersResponse.newBuilder() + .setResponse(CommonResponse.newBuilder().build()) + .build()) + .build()); + extProcLatch.countDown(); + } else if (request.hasResponseTrailers()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseTrailers( + TrailersResponse.newBuilder() + .setHeaderMutation( + HeaderMutation.newBuilder() + .addSetHeaders( + io.envoyproxy.envoy.config.core.v3.HeaderValueOption + .newBuilder() + .setHeader( + io.envoyproxy.envoy.config.core.v3.HeaderValue + .newBuilder() + .setKey("x-mutated-trailer") + .setValue("mutated-trailer-value") + .build()) + .build()) + .build()) + .build()) + .build()); + responseObserver.onCompleted(); + extProcLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final AtomicReference receivedTrailers = new AtomicReference<>(); + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + receivedTrailers.set(trailers); + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(receivedTrailers.get()).isNotNull(); + assertThat( + receivedTrailers + .get() + .get(Metadata.Key.of("x-mutated-trailer", Metadata.ASCII_STRING_MARSHALLER))) + .isEqualTo("mutated-trailer-value"); + } + + @Test + public void givenResponseTrailerModeSend_whenCallCloses_thenResponseTrailersSentToExtProc() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SEND) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = + new CountDownLatch(2); // 1 for request headers, 1 for response trailers + final AtomicReference capturedTrailerRequest = new AtomicReference<>(); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + extProcLatch.countDown(); + } else if (request.hasResponseTrailers()) { + capturedTrailerRequest.set(request); + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseTrailers(TrailersResponse.newBuilder().build()) + .build()); + responseObserver.onCompleted(); + extProcLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + ProcessingRequest req = capturedTrailerRequest.get(); + assertThat(req).isNotNull(); + assertThat(req.hasResponseTrailers()).isTrue(); + } + + @Test + public void givenResponseTrailerModeDefault_whenCallCloses_thenResponseTrailersNotSentToExtProc() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.DEFAULT) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = + new CountDownLatch(2); // 1 for request headers, 1 for response headers + final AtomicBoolean responseTrailersReceived = new AtomicBoolean(false); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + extProcLatch.countDown(); + } else if (request.hasResponseHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders(HeadersResponse.newBuilder().build()) + .build()); + responseObserver.onCompleted(); + extProcLatch.countDown(); + } else if (request.hasResponseTrailers()) { + responseTrailersReceived.set(true); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(responseTrailersReceived.get()).isFalse(); + } + + @Test + public void givenResponseTrailerModeSkip_whenCallCloses_thenResponseTrailersNotSentToExtProc() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SKIP) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = + new CountDownLatch(2); // 1 for request headers, 1 for response headers + final AtomicBoolean responseTrailersReceived = new AtomicBoolean(false); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + extProcLatch.countDown(); + } else if (request.hasResponseHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders(HeadersResponse.newBuilder().build()) + .build()); + responseObserver.onCompleted(); + extProcLatch.countDown(); + } else if (request.hasResponseTrailers()) { + responseTrailersReceived.set(true); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(responseTrailersReceived.get()).isFalse(); + } + + // ============================================================================ + // Category 11: Trailers-Only Response Handling + // ============================================================================ + + @Test + public void givenResponseHeaderModeSend_whenTrailersOnlySent_thenResponseHeadersSentToExtProc() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SEND) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = + new CountDownLatch(2); // 1 for request headers, 1 for response headers + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders( + HeadersResponse.newBuilder() + .setResponse(CommonResponse.newBuilder().build()) + .build()) + .build()); + extProcLatch.countDown(); + } else if (request.hasResponseHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders( + HeadersResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setHeaderMutation( + HeaderMutation.newBuilder() + .addSetHeaders( + io.envoyproxy.envoy.config.core.v3 + .HeaderValueOption.newBuilder() + .setHeader( + io.envoyproxy.envoy.config.core.v3 + .HeaderValue.newBuilder() + .setKey("x-mutated-trailer") + .setValue("mutated-trailer-value") + .build()) + .build()) + .build()) + .build()) + .build()) + .build()); + responseObserver.onCompleted(); + extProcLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + // Directly close with error without calling onNext / sendHeaders + responseObserver.onError( + Status.UNAUTHENTICATED + .withDescription("forced-trailers-only") + .asRuntimeException()); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final AtomicReference receivedTrailers = new AtomicReference<>(); + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + final AtomicReference receivedStatus = new AtomicReference<>(); + + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + receivedStatus.set(status); + receivedTrailers.set(trailers); + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(receivedStatus.get().getCode()).isEqualTo(Status.Code.UNAUTHENTICATED); + assertThat(receivedTrailers.get()).isNotNull(); + assertThat( + receivedTrailers + .get() + .get(Metadata.Key.of("x-mutated-trailer", Metadata.ASCII_STRING_MARSHALLER))) + .isEqualTo("mutated-trailer-value"); + } + + @Test + public void givenResponseHeaderModeDefault_whenTrailersOnlySent_thenResponseHeadersSentToExtProc() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.DEFAULT) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = + new CountDownLatch(2); // 1 for request headers, 1 for response headers + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders( + HeadersResponse.newBuilder() + .setResponse(CommonResponse.newBuilder().build()) + .build()) + .build()); + extProcLatch.countDown(); + } else if (request.hasResponseHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders( + HeadersResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setHeaderMutation( + HeaderMutation.newBuilder() + .addSetHeaders( + io.envoyproxy.envoy.config.core.v3 + .HeaderValueOption.newBuilder() + .setHeader( + io.envoyproxy.envoy.config.core.v3 + .HeaderValue.newBuilder() + .setKey("x-mutated-trailer") + .setValue("mutated-trailer-value") + .build()) + .build()) + .build()) + .build()) + .build()) + .build()); + responseObserver.onCompleted(); + extProcLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + // Directly close with error without calling onNext / sendHeaders + responseObserver.onError( + Status.UNAUTHENTICATED + .withDescription("forced-trailers-only") + .asRuntimeException()); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final AtomicReference receivedTrailers = new AtomicReference<>(); + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + final AtomicReference receivedStatus = new AtomicReference<>(); + + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + receivedStatus.set(status); + receivedTrailers.set(trailers); + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(receivedStatus.get().getCode()).isEqualTo(Status.Code.UNAUTHENTICATED); + assertThat(receivedTrailers.get()).isNotNull(); + assertThat( + receivedTrailers + .get() + .get(Metadata.Key.of("x-mutated-trailer", Metadata.ASCII_STRING_MARSHALLER))) + .isEqualTo("mutated-trailer-value"); + } + + @Test + public void givenResponseHeaderModeSkip_whenTrailersOnlySent_thenResponseHeadersNotSentToExtProc() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = + new CountDownLatch(1); // 1 for request headers, response headers/trailers skipped + final AtomicBoolean responseHeadersReceived = new AtomicBoolean(false); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + extProcLatch.countDown(); + } else if (request.hasResponseHeaders()) { + responseHeadersReceived.set(true); + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders(HeadersResponse.newBuilder().build()) + .build()); + responseObserver.onCompleted(); + extProcLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + Metadata trailers = new Metadata(); + trailers.put( + Metadata.Key.of("x-dataplane-trailer", Metadata.ASCII_STRING_MARSHALLER), + "original"); + responseObserver.onError( + Status.UNAUTHENTICATED + .withDescription("forced-trailers-only") + .asRuntimeException(trailers)); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final AtomicReference receivedTrailers = new AtomicReference<>(); + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + final AtomicReference receivedStatus = new AtomicReference<>(); + + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + receivedStatus.set(status); + receivedTrailers.set(trailers); + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(responseHeadersReceived.get()).isFalse(); + assertThat(receivedStatus.get().getCode()).isEqualTo(Status.Code.UNAUTHENTICATED); + assertThat(receivedStatus.get().getDescription()).isEqualTo("forced-trailers-only"); + assertThat(receivedTrailers.get()).isNotNull(); + assertThat( + receivedTrailers + .get() + .get(Metadata.Key.of("x-dataplane-trailer", Metadata.ASCII_STRING_MARSHALLER))) + .isEqualTo("original"); + } + + // ============================================================================ + // Category 12: Half-Close handling + // ============================================================================ + + @Test + public void givenRequestBodyModeGrpc_whenHalfCloseCalled_extProcCalledWithEos() throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestBodyMode(ProcessingMode.BodySendMode.GRPC) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = new CountDownLatch(2); // 1 for headers, 1 for body EOS + final AtomicBoolean receivedEos = new AtomicBoolean(false); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders( + HeadersResponse.newBuilder() + .setResponse(CommonResponse.newBuilder().build()) + .build()) + .build()); + extProcLatch.countDown(); + } else if (request.hasRequestBody()) { + if (request.getRequestBody().getEndOfStreamWithoutMessage()) { + receivedEos.set(true); + extProcLatch.countDown(); + } else { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestBody( + BodyResponse.newBuilder() + .setResponse(CommonResponse.newBuilder().build()) + .build()) + .build()); + } + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final CountDownLatch appHalfCloseLatch = new CountDownLatch(1); + final List serverReceivedMessages = new java.util.concurrent.CopyOnWriteArrayList<>(); + ServerInterceptor capturingInterceptor = + new ServerInterceptor() { + @Override + public ServerCall.Listener interceptCall( + ServerCall call, Metadata headers, ServerCallHandler next) { + ServerCall.Listener nextListener = next.startCall(call, headers); + return new io.grpc.ForwardingServerCallListener.SimpleForwardingServerCallListener< + ReqT>(nextListener) { + @Override + @SuppressWarnings("unchecked") + public void onMessage(ReqT message) { + try { + byte[] bytes = ByteStreams.toByteArray((InputStream) message); + serverReceivedMessages.add(new String(bytes, StandardCharsets.UTF_8)); + super.onMessage((ReqT) new ByteArrayInputStream(bytes)); + } catch (Exception e) { + throw new RuntimeException(e); + } + } + + @Override + public void onHalfClose() { + appHalfCloseLatch.countDown(); + super.onHalfClose(); + } + }; + } + }; + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext( + new ByteArrayInputStream("response".getBytes(StandardCharsets.UTF_8))); + responseObserver.onCompleted(); + } + }; + + startDataPlane(capturingInterceptor, interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + try { + clientCall.start(new io.grpc.ClientCall.Listener() {}, new Metadata()); + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertWithMessage("Ext proc should receive headers and body EOS") + .that(extProcLatch.await(5, TimeUnit.SECONDS)) + .isTrue(); + assertThat(receivedEos.get()).isTrue(); + assertThat(appHalfCloseLatch.await(500, TimeUnit.MILLISECONDS)).isFalse(); + assertThat(serverReceivedMessages).isEmpty(); + } finally { + clientCall.cancel("Cleanup", null); + channelManager.close(); + } + } + + @Test + public void + deferredHalfClose_whenExtProcRespondsWithEosWithoutMessage_thenAppListenerReceivesHalfClose() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestBodyMode(ProcessingMode.BodySendMode.GRPC) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = new CountDownLatch(2); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders( + HeadersResponse.newBuilder() + .setResponse(CommonResponse.newBuilder().build()) + .build()) + .build()); + extProcLatch.countDown(); + } else if (request.hasRequestBody()) { + if (request.getRequestBody().getEndOfStreamWithoutMessage()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestBody( + BodyResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setEndOfStreamWithoutMessage(true) + .build()) + .build()) + .build()) + .build()) + .build()); + extProcLatch.countDown(); + } else { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestBody( + BodyResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setBody( + request.getRequestBody().getBody()) + .build()) + .build()) + .build()) + .build()) + .build()); + } + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final CountDownLatch appHalfCloseLatch = new CountDownLatch(1); + final List serverReceivedMessages = new java.util.concurrent.CopyOnWriteArrayList<>(); + ServerInterceptor capturingInterceptor = + new ServerInterceptor() { + @Override + public ServerCall.Listener interceptCall( + ServerCall call, Metadata headers, ServerCallHandler next) { + ServerCall.Listener nextListener = next.startCall(call, headers); + return new io.grpc.ForwardingServerCallListener.SimpleForwardingServerCallListener< + ReqT>(nextListener) { + @Override + @SuppressWarnings("unchecked") + public void onMessage(ReqT message) { + try { + byte[] bytes = ByteStreams.toByteArray((InputStream) message); + serverReceivedMessages.add(new String(bytes, StandardCharsets.UTF_8)); + super.onMessage((ReqT) new ByteArrayInputStream(bytes)); + } catch (Exception e) { + throw new RuntimeException(e); + } + } + + @Override + public void onHalfClose() { + appHalfCloseLatch.countDown(); + super.onHalfClose(); + } + }; + } + }; + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext( + new ByteArrayInputStream("response".getBytes(StandardCharsets.UTF_8))); + responseObserver.onCompleted(); + } + }; + + startDataPlane(capturingInterceptor, interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + try { + clientCall.start(new io.grpc.ClientCall.Listener() {}, new Metadata()); + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertWithMessage("Ext proc should receive headers and body EOS") + .that(extProcLatch.await(5, TimeUnit.SECONDS)) + .isTrue(); + assertThat(appHalfCloseLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(serverReceivedMessages).containsExactly("hello"); + } finally { + clientCall.cancel("Cleanup", null); + channelManager.close(); + } + } + + @Test + public void + givenDeferredHalfClose_whenExtProcRespondsWithEndOfStream_thenAppListenerReceivesHalfClose() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestBodyMode(ProcessingMode.BodySendMode.GRPC) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = new CountDownLatch(2); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders( + HeadersResponse.newBuilder() + .setResponse(CommonResponse.newBuilder().build()) + .build()) + .build()); + extProcLatch.countDown(); + } else if (request.hasRequestBody()) { + if (request.getRequestBody().getEndOfStreamWithoutMessage()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestBody( + BodyResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setBody( + ByteString.copyFromUtf8( + " mutated-eof")) + .setEndOfStream(true) + .build()) + .build()) + .build()) + .build()) + .build()); + extProcLatch.countDown(); + } else { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestBody( + BodyResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setBody( + request.getRequestBody().getBody()) + .build()) + .build()) + .build()) + .build()) + .build()); + } + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final CountDownLatch appHalfCloseLatch = new CountDownLatch(1); + final List serverReceivedMessages = new java.util.concurrent.CopyOnWriteArrayList<>(); + ServerInterceptor capturingInterceptor = + new ServerInterceptor() { + @Override + public ServerCall.Listener interceptCall( + ServerCall call, Metadata headers, ServerCallHandler next) { + ServerCall.Listener nextListener = next.startCall(call, headers); + return new io.grpc.ForwardingServerCallListener.SimpleForwardingServerCallListener< + ReqT>(nextListener) { + @Override + @SuppressWarnings("unchecked") + public void onMessage(ReqT message) { + try { + byte[] bytes = ByteStreams.toByteArray((InputStream) message); + serverReceivedMessages.add(new String(bytes, StandardCharsets.UTF_8)); + super.onMessage((ReqT) new ByteArrayInputStream(bytes)); + } catch (Exception e) { + throw new RuntimeException(e); + } + } + + @Override + public void onHalfClose() { + appHalfCloseLatch.countDown(); + super.onHalfClose(); + } + }; + } + }; + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext( + new ByteArrayInputStream("response".getBytes(StandardCharsets.UTF_8))); + responseObserver.onCompleted(); + } + }; + + startDataPlane(capturingInterceptor, interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + try { + clientCall.start(new io.grpc.ClientCall.Listener() {}, new Metadata()); + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertWithMessage("Ext proc should receive headers and body EOS") + .that(extProcLatch.await(5, TimeUnit.SECONDS)) + .isTrue(); + assertThat(appHalfCloseLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(serverReceivedMessages).containsExactly("hello", " mutated-eof"); + } finally { + clientCall.cancel("Cleanup", null); + channelManager.close(); + } + } + + @Test + public void extProcEosNoMsg_whenClientNotHalfClosed_thenAppHalfClosed_moreMessagesDiscarded() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestBodyMode(ProcessingMode.BodySendMode.GRPC) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final List extProcReceivedBodies = + new java.util.concurrent.CopyOnWriteArrayList<>(); + final CountDownLatch extProcLatch = new CountDownLatch(2); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders( + HeadersResponse.newBuilder() + .setResponse(CommonResponse.newBuilder().build()) + .build()) + .build()); + extProcLatch.countDown(); + } else if (request.hasRequestBody()) { + extProcReceivedBodies.add(request.getRequestBody().getBody()); + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestBody( + BodyResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setEndOfStreamWithoutMessage(true) + .build()) + .build()) + .build()) + .build()) + .build()); + extProcLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final CountDownLatch appHalfCloseLatch = new CountDownLatch(1); + final List serverReceivedMessages = new java.util.concurrent.CopyOnWriteArrayList<>(); + ServerInterceptor capturingInterceptor = + new ServerInterceptor() { + @Override + public ServerCall.Listener interceptCall( + ServerCall call, Metadata headers, ServerCallHandler next) { + ServerCall.Listener nextListener = next.startCall(call, headers); + return new io.grpc.ForwardingServerCallListener.SimpleForwardingServerCallListener< + ReqT>(nextListener) { + @Override + @SuppressWarnings("unchecked") + public void onMessage(ReqT message) { + try { + byte[] bytes = ByteStreams.toByteArray((InputStream) message); + serverReceivedMessages.add(new String(bytes, StandardCharsets.UTF_8)); + super.onMessage((ReqT) new ByteArrayInputStream(bytes)); + } catch (Exception e) { + throw new RuntimeException(e); + } + } + + @Override + public void onHalfClose() { + appHalfCloseLatch.countDown(); + super.onHalfClose(); + } + }; + } + }; + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext( + new ByteArrayInputStream("response".getBytes(StandardCharsets.UTF_8))); + responseObserver.onCompleted(); + } + }; + + startDataPlane(capturingInterceptor, interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + try { + clientCall.start(new io.grpc.ClientCall.Listener() {}, new Metadata()); + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + + assertWithMessage("Ext proc should receive request headers and body") + .that(extProcLatch.await(5, TimeUnit.SECONDS)) + .isTrue(); + assertThat(appHalfCloseLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + // Client sends another message after half-close has been propagated to app listener + clientCall.sendMessage( + new ByteArrayInputStream("extra-message".getBytes(StandardCharsets.UTF_8))); + + // Wait a short time to ensure it is discarded + Thread.sleep(200); + + // Verify app listener received no messages (discarded due to EOS without message) and not + // "extra-message" + assertThat(serverReceivedMessages).isEmpty(); + // Verify ext-proc did not receive "extra-message" either + assertThat(extProcReceivedBodies).containsExactly(ByteString.copyFromUtf8("hello")); + + clientCall.halfClose(); + } finally { + clientCall.cancel("Cleanup", null); + channelManager.close(); + } + } + + @Test + public void extProcEos_whenClientNotHalfClosed_thenAppHalfClosed_moreMessagesDiscarded() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestBodyMode(ProcessingMode.BodySendMode.GRPC) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final List extProcReceivedBodies = + new java.util.concurrent.CopyOnWriteArrayList<>(); + final CountDownLatch extProcLatch = new CountDownLatch(2); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders( + HeadersResponse.newBuilder() + .setResponse(CommonResponse.newBuilder().build()) + .build()) + .build()); + extProcLatch.countDown(); + } else if (request.hasRequestBody()) { + extProcReceivedBodies.add(request.getRequestBody().getBody()); + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestBody( + BodyResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setBody( + ByteString.copyFromUtf8( + "hello mutated")) + .setEndOfStream(true) + .build()) + .build()) + .build()) + .build()) + .build()); + extProcLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final CountDownLatch appHalfCloseLatch = new CountDownLatch(1); + final List serverReceivedMessages = new java.util.concurrent.CopyOnWriteArrayList<>(); + ServerInterceptor capturingInterceptor = + new ServerInterceptor() { + @Override + public ServerCall.Listener interceptCall( + ServerCall call, Metadata headers, ServerCallHandler next) { + ServerCall.Listener nextListener = next.startCall(call, headers); + return new io.grpc.ForwardingServerCallListener.SimpleForwardingServerCallListener< + ReqT>(nextListener) { + @Override + @SuppressWarnings("unchecked") + public void onMessage(ReqT message) { + try { + byte[] bytes = ByteStreams.toByteArray((InputStream) message); + serverReceivedMessages.add(new String(bytes, StandardCharsets.UTF_8)); + super.onMessage((ReqT) new ByteArrayInputStream(bytes)); + } catch (Exception e) { + throw new RuntimeException(e); + } + } + + @Override + public void onHalfClose() { + appHalfCloseLatch.countDown(); + super.onHalfClose(); + } + }; + } + }; + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext( + new ByteArrayInputStream("response".getBytes(StandardCharsets.UTF_8))); + responseObserver.onCompleted(); + } + }; + + startDataPlane(capturingInterceptor, interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + try { + clientCall.start(new io.grpc.ClientCall.Listener() {}, new Metadata()); + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + + assertWithMessage("Ext proc should receive request headers and body") + .that(extProcLatch.await(5, TimeUnit.SECONDS)) + .isTrue(); + assertThat(appHalfCloseLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(serverReceivedMessages).containsExactly("hello mutated"); + + // Client sends another message after half-close has been propagated to app listener + clientCall.sendMessage( + new ByteArrayInputStream("extra-message".getBytes(StandardCharsets.UTF_8))); + + // Wait a short time to ensure it is discarded + Thread.sleep(200); + + // Verify app listener only received "hello mutated" and not "extra-message" + assertThat(serverReceivedMessages).containsExactly("hello mutated"); + // Verify ext-proc did not receive "extra-message" either + assertThat(extProcReceivedBodies).containsExactly(ByteString.copyFromUtf8("hello")); + + clientCall.halfClose(); + } finally { + clientCall.cancel("Cleanup", null); + channelManager.close(); + } + } + + // ============================================================================ + // Category 13: Outbound Backpressure (isReady / onReady) + // ============================================================================ + + @Test + @SuppressWarnings("unchecked") + public void givenObservabilityTrue_whenExtProcBusy_thenIsReadyReturnsFalse() throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + ExternalProcessor proto = + ExternalProcessor.newBuilder() + .setGrpcService( + GrpcService.newBuilder() + .setGoogleGrpc( + GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///" + uniqueExtProcServerName) + .addChannelCredentialsPlugin( + Any.newBuilder() + .setTypeUrl( + "type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build()) + .setObservabilityMode(true) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + // External Processor Server + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + responseObserverRef.set(responseObserver); + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + final AtomicBoolean sidecarReady = new AtomicBoolean(true); + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName) + .directExecutor() + .intercept( + new io.grpc.ClientInterceptor() { + @Override + public io.grpc.ClientCall interceptCall( + MethodDescriptor method, + io.grpc.CallOptions callOptions, + io.grpc.Channel next) { + return new io.grpc.ForwardingClientCall.SimpleForwardingClientCall< + ReqT, RespT>(next.newCall(method, callOptions)) { + @Override + public boolean isReady() { + return sidecarReady.get(); + } + }; + } + }) + .build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicReference> interceptedCallRef = + new AtomicReference<>(); + ServerCallHandler nextHandler = + (call, headers) -> { + interceptedCallRef.set(call); + return new ServerCall.Listener() {}; + }; + + final AtomicBoolean rawCallReady = new AtomicBoolean(true); + ServerCall rawCall = + new SimpleServerCall(METHOD_SAY_HELLO_RAW) { + @Override + public boolean isReady() { + return rawCallReady.get(); + } + }; + + try { + interceptor.interceptCall(rawCall, new Metadata(), nextHandler); + + // Wait for activation (ext_proc response received) + long startTime = System.currentTimeMillis(); + while (!interceptedCallRef.get().isReady() && System.currentTimeMillis() - startTime < 5000) { + Thread.sleep(10); + } + assertThat(interceptedCallRef.get().isReady()).isTrue(); + + // Ext proc busy + sidecarReady.set(false); + assertThat(interceptedCallRef.get().isReady()).isFalse(); + + // Ext proc ready again + sidecarReady.set(true); + assertThat(interceptedCallRef.get().isReady()).isTrue(); + + // Raw call not ready + rawCallReady.set(false); + assertThat(interceptedCallRef.get().isReady()).isFalse(); + } finally { + if (responseObserverRef.get() != null) { + try { + responseObserverRef.get().onCompleted(); + } catch (IllegalStateException ignored) { + // Ignore if already closed + } + } + channelManager.close(); + } + } + + @Test + @SuppressWarnings("unchecked") + public void givenObservabilityMode_whenClientBusy_thenIsReadyReturnsFalse() throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + ExternalProcessor proto = + ExternalProcessor.newBuilder() + .setGrpcService( + GrpcService.newBuilder() + .setGoogleGrpc( + GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///" + uniqueExtProcServerName) + .addChannelCredentialsPlugin( + Any.newBuilder() + .setTypeUrl( + "type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build()) + .setObservabilityMode(true) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + // External Processor Server + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + responseObserverRef.set(responseObserver); + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName) + .directExecutor() + .build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicBoolean clientReady = new AtomicBoolean(true); + ServerCall rawCall = + new SimpleServerCall(METHOD_SAY_HELLO_RAW) { + @Override + public boolean isReady() { + return clientReady.get(); + } + }; + + final AtomicReference> interceptedCallRef = + new AtomicReference<>(); + try { + interceptor.interceptCall( + rawCall, + new Metadata(), + (call, headers) -> { + interceptedCallRef.set(call); + return new ServerCall.Listener() {}; + }); + + // Wait for activation (sidecar needs to respond to headers) + long startTime = System.currentTimeMillis(); + while (interceptedCallRef.get() == null && System.currentTimeMillis() - startTime < 5000) { + Thread.sleep(10); + } + assertThat(interceptedCallRef.get()).isNotNull(); + assertThat(interceptedCallRef.get().isReady()).isTrue(); + + // Client becomes busy + clientReady.set(false); + assertThat(interceptedCallRef.get().isReady()).isFalse(); + } finally { + if (responseObserverRef.get() != null) { + try { + responseObserverRef.get().onCompleted(); + } catch (IllegalStateException ignored) { + // ignore + } + } + channelManager.close(); + } + } + + @Test + @SuppressWarnings("unchecked") + public void givenNormalMode_whenClientBusy_thenIsReadyReturnsTrue() throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + ExternalProcessor proto = + ExternalProcessor.newBuilder() + .setGrpcService( + GrpcService.newBuilder() + .setGoogleGrpc( + GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///" + uniqueExtProcServerName) + .addChannelCredentialsPlugin( + Any.newBuilder() + .setTypeUrl( + "type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build()) + .setObservabilityMode(false) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + // External Processor Server + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + responseObserverRef.set(responseObserver); + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName) + .directExecutor() + .build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicBoolean clientReady = new AtomicBoolean(false); + ServerCall rawCall = + new SimpleServerCall(METHOD_SAY_HELLO_RAW) { + @Override + public boolean isReady() { + return clientReady.get(); + } + }; + + final AtomicReference> interceptedCallRef = + new AtomicReference<>(); + try { + interceptor.interceptCall( + rawCall, + new Metadata(), + (call, headers) -> { + interceptedCallRef.set(call); + return new ServerCall.Listener() {}; + }); + + // Wait for activation (sidecar needs to respond to headers) + long startTime = System.currentTimeMillis(); + while (interceptedCallRef.get() == null && System.currentTimeMillis() - startTime < 5000) { + Thread.sleep(10); + } + assertThat(interceptedCallRef.get()).isNotNull(); + + // Since sidecar is ready, interceptedCallRef.get().isReady() should return true, + // ignoring that client (rawCall) is busy + assertThat(interceptedCallRef.get().isReady()).isTrue(); + } finally { + if (responseObserverRef.get() != null) { + try { + responseObserverRef.get().onCompleted(); + } catch (IllegalStateException ignored) { + // ignore + } + } + channelManager.close(); + } + } + + @Test + @SuppressWarnings("unchecked") + public void givenCongestionInExtProc_whenExtProcBecomesReady_thenTriggersOnReady() + throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + ExternalProcessor proto = + ExternalProcessor.newBuilder() + .setGrpcService( + GrpcService.newBuilder() + .setGoogleGrpc( + GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///" + uniqueExtProcServerName) + .addChannelCredentialsPlugin( + Any.newBuilder() + .setTypeUrl( + "type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build()) + .setObservabilityMode(true) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + // External Processor Server + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + responseObserverRef.set(responseObserver); + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) {} + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + final AtomicReference> sidecarListenerRef = + new AtomicReference<>(); + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName) + .directExecutor() + .intercept( + new io.grpc.ClientInterceptor() { + @Override + public io.grpc.ClientCall interceptCall( + MethodDescriptor method, + io.grpc.CallOptions callOptions, + io.grpc.Channel next) { + return new io.grpc.ForwardingClientCall.SimpleForwardingClientCall< + ReqT, RespT>(next.newCall(method, callOptions)) { + @Override + public void start( + Listener responseListener, Metadata headers) { + sidecarListenerRef.set( + (Listener) responseListener); + super.start(responseListener, headers); + } + }; + } + }) + .build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final CountDownLatch onReadyLatch = new CountDownLatch(1); + ServerCall.Listener appListener = + new ServerCall.Listener() { + @Override + public void onReady() { + onReadyLatch.countDown(); + } + }; + + ServerCallHandler nextHandler = + (call, headers) -> { + return appListener; + }; + + try { + interceptor.interceptCall( + new SimpleServerCall(METHOD_SAY_HELLO_RAW), new Metadata(), nextHandler); + + // Wait for sidecar call to start and listener to be captured + long startTime = System.currentTimeMillis(); + while (sidecarListenerRef.get() == null && System.currentTimeMillis() - startTime < 5000) { + Thread.sleep(10); + } + assertThat(sidecarListenerRef.get()).isNotNull(); + + // Trigger sidecar onReady + sidecarListenerRef.get().onReady(); + + // Verify app listener notified + assertThat(onReadyLatch.await(5, TimeUnit.SECONDS)).isTrue(); + } finally { + if (responseObserverRef.get() != null) { + try { + responseObserverRef.get().onCompleted(); + } catch (IllegalStateException ignored) { + // Ignore if already closed + } + } + channelManager.close(); + } + } + + // ============================================================================ + // Category 14: Ext-proc request draining + // ============================================================================ + + @Test + @SuppressWarnings("unchecked") + public void givenRequestDrainActive_whenIsReadyCalled_thenReturnsFalse() throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + ExternalProcessor proto = + ExternalProcessor.newBuilder() + .setGrpcService( + GrpcService.newBuilder() + .setGoogleGrpc( + GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///" + uniqueExtProcServerName) + .addChannelCredentialsPlugin( + Any.newBuilder() + .setTypeUrl( + "type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch drainLatch = new CountDownLatch(1); + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + responseObserverRef.set(responseObserver); + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder().setRequestDrain(true).build()); + drainLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName) + .directExecutor() + .build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicReference> interceptedCallRef = + new AtomicReference<>(); + ServerCallHandler nextHandler = + (call, headers) -> { + interceptedCallRef.set(call); + return new ServerCall.Listener() {}; + }; + + try { + interceptor.interceptCall( + new SimpleServerCall(METHOD_SAY_HELLO_RAW), new Metadata(), nextHandler); + + assertThat(drainLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + // isReady() must return false during drain. + long start = System.currentTimeMillis(); + while (interceptedCallRef.get().isReady() && System.currentTimeMillis() - start < 2000) { + Thread.sleep(10); + } + assertThat(interceptedCallRef.get().isReady()).isFalse(); + } finally { + if (responseObserverRef.get() != null) { + try { + responseObserverRef.get().onCompleted(); + } catch (IllegalStateException ignored) { + // Ignore if already closed + } + } + channelManager.close(); + } + } + + @Test + @SuppressWarnings("unchecked") + public void givenDrainingStream_whenExtProcStreamCompletes_thenOnReady() throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + ExternalProcessor proto = + ExternalProcessor.newBuilder() + .setGrpcService( + GrpcService.newBuilder() + .setGoogleGrpc( + GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///" + uniqueExtProcServerName) + .addChannelCredentialsPlugin( + Any.newBuilder() + .setTypeUrl( + "type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch sidecarFinishLatch = new CountDownLatch(1); + final CountDownLatch sidecarOnNextLatch = new CountDownLatch(1); + final CountDownLatch sidecarOnCompletedLatch = new CountDownLatch(1); + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + responseObserverRef.set(responseObserver); + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + new Thread( + () -> { + responseObserver.onNext( + ProcessingResponse.newBuilder().setRequestDrain(true).build()); + sidecarOnNextLatch.countDown(); + try { + if (sidecarFinishLatch.await(5, TimeUnit.SECONDS)) { + sidecarOnCompletedLatch.countDown(); + responseObserver.onCompleted(); + } + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + }) + .start(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName) + .directExecutor() + .build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final CountDownLatch onReadyLatch = new CountDownLatch(1); + ServerCall.Listener appListener = + new ServerCall.Listener() { + @Override + public void onReady() { + onReadyLatch.countDown(); + } + }; + + final AtomicReference> interceptedCallRef = + new AtomicReference<>(); + ServerCall rawCall = + new SimpleServerCall(METHOD_SAY_HELLO_RAW) { + @Override + public boolean isReady() { + return true; + } + }; + + try { + interceptor.interceptCall( + rawCall, + new Metadata(), + (call, headers) -> { + interceptedCallRef.set(call); + return appListener; + }); + + // Wait for sidecar to send drain and test to observe it + assertThat(sidecarOnNextLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + long start = System.currentTimeMillis(); + while (interceptedCallRef.get() == null && System.currentTimeMillis() - start < 2000) { + Thread.sleep(10); + } + assertThat(interceptedCallRef.get()).isNotNull(); + assertThat(interceptedCallRef.get().isReady()).isFalse(); + + // Now let sidecar complete + sidecarFinishLatch.countDown(); + + assertThat(sidecarOnCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + // After sidecar stream completes, it should trigger onReady and become ready + assertThat(onReadyLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(interceptedCallRef.get().isReady()).isTrue(); + } finally { + if (responseObserverRef.get() != null) { + try { + responseObserverRef.get().onCompleted(); + } catch (IllegalStateException ignored) { + // ignore + } + } + channelManager.close(); + } + } + + @Test + @SuppressWarnings("unchecked") + public void givenDrainingStream_whenExtProcStreamCompletes_thenMessagesProceed() + throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + ExternalProcessor proto = + ExternalProcessor.newBuilder() + .setGrpcService( + GrpcService.newBuilder() + .setGoogleGrpc( + GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///" + uniqueExtProcServerName) + .addChannelCredentialsPlugin( + Any.newBuilder() + .setTypeUrl( + "type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + // External Processor Server + final CountDownLatch sidecarFinishLatch = new CountDownLatch(1); + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + responseObserverRef.set(responseObserver); + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + new Thread( + () -> { + responseObserver.onNext( + ProcessingResponse.newBuilder().setRequestDrain(true).build()); + try { + if (sidecarFinishLatch.await(5, TimeUnit.SECONDS)) { + responseObserver.onCompleted(); + } + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + }) + .start(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName) + .directExecutor() + .build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final List serverReceivedMessages = new java.util.concurrent.CopyOnWriteArrayList<>(); + final CountDownLatch appMessageLatch = new CountDownLatch(1); + ServerCall.Listener appListener = + new ServerCall.Listener() { + @Override + public void onMessage(InputStream message) { + try { + serverReceivedMessages.add( + new String(ByteStreams.toByteArray(message), StandardCharsets.UTF_8)); + appMessageLatch.countDown(); + } catch (IOException e) { + throw new RuntimeException(e); + } + } + }; + + final AtomicReference> interceptedCallRef = + new AtomicReference<>(); + final List rawSentMessages = new java.util.concurrent.CopyOnWriteArrayList<>(); + final CountDownLatch rawSentLatch = new CountDownLatch(1); + ServerCall rawCall = + new SimpleServerCall(METHOD_SAY_HELLO_RAW) { + @Override + public void sendMessage(InputStream message) { + rawSentMessages.add(message); + rawSentLatch.countDown(); + } + + @Override + public boolean isReady() { + return true; + } + }; + + try { + ServerCall.Listener interceptedListener = + interceptor.interceptCall( + rawCall, + new Metadata(), + (call, headers) -> { + interceptedCallRef.set(call); + return appListener; + }); + + // Wait for drain to be processed + long start = System.currentTimeMillis(); + while ((interceptedCallRef.get() == null || interceptedCallRef.get().isReady()) + && System.currentTimeMillis() - start < 5000) { + Thread.sleep(10); + } + assertThat(interceptedCallRef.get().isReady()).isFalse(); + + // Now let sidecar complete + sidecarFinishLatch.countDown(); + + // Wait for it to become ready again + start = System.currentTimeMillis(); + while (!interceptedCallRef.get().isReady() && System.currentTimeMillis() - start < 5000) { + Thread.sleep(10); + } + assertThat(interceptedCallRef.get().isReady()).isTrue(); + + // 1. Verify client message is delivered to app listener without sidecar contact + interceptedListener.onMessage( + new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + assertThat(appMessageLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(serverReceivedMessages).containsExactly("hello"); + + // 2. Verify server app response message is sent to client without sidecar contact + interceptedCallRef + .get() + .sendMessage(new ByteArrayInputStream("response".getBytes(StandardCharsets.UTF_8))); + assertThat(rawSentLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat( + new String(ByteStreams.toByteArray(rawSentMessages.get(0)), StandardCharsets.UTF_8)) + .isEqualTo("response"); + } finally { + if (responseObserverRef.get() != null) { + try { + responseObserverRef.get().onCompleted(); + } catch (IllegalStateException ignored) { + // ignore + } + } + channelManager.close(); + } + } + + @Test + @SuppressWarnings("unchecked") + public void + drainingStartsBeforeResponseHeaders_whenAppSendsMessagesAndStatus_thenBufferedAndDelivered() + throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + ExternalProcessor proto = + ExternalProcessor.newBuilder() + .setGrpcService( + GrpcService.newBuilder() + .setGoogleGrpc( + GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///" + uniqueExtProcServerName) + .addChannelCredentialsPlugin( + Any.newBuilder() + .setTypeUrl( + "type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build()) + .setProcessingMode( + ProcessingMode.newBuilder() + .setResponseBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SEND) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + // External Processor Server + final CountDownLatch sidecarFinishLatch = new CountDownLatch(1); + final CountDownLatch drainCompletedLatch = new CountDownLatch(1); + final AtomicInteger extProcReceivedBodyCount = new AtomicInteger(0); + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + responseObserverRef.set(responseObserver); + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + new Thread( + () -> { + responseObserver.onNext( + ProcessingResponse.newBuilder().setRequestDrain(true).build()); + try { + if (sidecarFinishLatch.await(5, TimeUnit.SECONDS)) { + responseObserver.onCompleted(); + } + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + }) + .start(); + } else if (request.hasResponseBody()) { + extProcReceivedBodyCount.incrementAndGet(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + drainCompletedLatch.countDown(); + } + }; + } + }; + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName) + .directExecutor() + .build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicReference> interceptedCallRef = + new AtomicReference<>(); + final List rawSentHeaders = new java.util.concurrent.CopyOnWriteArrayList<>(); + final List rawSentMessages = new java.util.concurrent.CopyOnWriteArrayList<>(); + final AtomicReference rawSentStatus = new AtomicReference<>(); + final AtomicReference rawSentTrailers = new AtomicReference<>(); + final CountDownLatch rawCloseLatch = new CountDownLatch(1); + + ServerCall rawCall = + new SimpleServerCall(METHOD_SAY_HELLO_RAW) { + @Override + public void sendHeaders(Metadata headers) { + rawSentHeaders.add(headers); + } + + @Override + public void sendMessage(InputStream message) { + rawSentMessages.add(message); + } + + @Override + public void close(Status status, Metadata trailers) { + rawSentStatus.set(status); + rawSentTrailers.set(trailers); + rawCloseLatch.countDown(); + } + + @Override + public boolean isReady() { + return true; + } + }; + + try { + interceptor.interceptCall( + rawCall, + new Metadata(), + (call, headers) -> { + interceptedCallRef.set(call); + return new ServerCall.Listener() {}; + }); + + // Wait for drain to be processed + assertThat(drainCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(interceptedCallRef.get().isReady()).isFalse(); + + // App sends response headers, message and closes server-side concurrently during drain + final CountDownLatch appActionLatch = new CountDownLatch(1); + new Thread( + () -> { + ServerCall interceptedCall = interceptedCallRef.get(); + interceptedCall.sendHeaders(new Metadata()); + interceptedCall.sendMessage( + new ByteArrayInputStream( + "response message during drain".getBytes(StandardCharsets.UTF_8))); + interceptedCall.close(Status.OK, new Metadata()); + appActionLatch.countDown(); + }) + .start(); + + assertThat(appActionLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + // Assert that it was NOT received by extProc + assertThat(extProcReceivedBodyCount.get()).isEqualTo(0); + // Assert that nothing has been delivered to client (rawCall) yet because drain is active + assertThat(rawSentHeaders).isEmpty(); + assertThat(rawSentMessages).isEmpty(); + assertThat(rawSentStatus.get()).isNull(); + + // Now let sidecar complete the drain + sidecarFinishLatch.countDown(); + + // Wait for rawCall to close + assertThat(rawCloseLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + // Verify delivery order: headers first, then app response message during drain, then close + assertThat(rawSentHeaders).hasSize(1); + + List deliveredMessages = new ArrayList<>(); + for (InputStream is : rawSentMessages) { + deliveredMessages.add(new String(ByteStreams.toByteArray(is), StandardCharsets.UTF_8)); + } + assertThat(deliveredMessages).containsExactly("response message during drain").inOrder(); + + assertThat(rawSentStatus.get().isOk()).isTrue(); + assertThat(rawSentTrailers.get()).isNotNull(); + } finally { + if (responseObserverRef.get() != null) { + try { + responseObserverRef.get().onCompleted(); + } catch (IllegalStateException ignored) { + // ignore + } + } + channelManager.close(); + } + } + + @Test + @SuppressWarnings("unchecked") + public void + drainingStartsAfterResponseHeaders_whenAppSendsMessagesAndStatus_thenBufferedAndDelivered() + throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + ExternalProcessor proto = + ExternalProcessor.newBuilder() + .setGrpcService( + GrpcService.newBuilder() + .setGoogleGrpc( + GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///" + uniqueExtProcServerName) + .addChannelCredentialsPlugin( + Any.newBuilder() + .setTypeUrl( + "type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build()) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SEND) + .setResponseBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SEND) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + // External Processor Server + final CountDownLatch sidecarFinishLatch = new CountDownLatch(1); + final CountDownLatch drainCompletedLatch = new CountDownLatch(1); + final AtomicInteger extProcReceivedBodyCount = new AtomicInteger(0); + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + responseObserverRef.set(responseObserver); + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasResponseHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders(HeadersResponse.newBuilder().build()) + .build()); + } else if (request.hasResponseBody()) { + int bodyIdx = extProcReceivedBodyCount.incrementAndGet(); + if (bodyIdx == 1) { + // Send mutated response body for the first original message + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseBody( + BodyResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setBody( + ByteString.copyFromUtf8( + "mutated-msg-1")) + .build()) + .build()) + .build()) + .build()) + .build()); + } else if (bodyIdx == 2) { + // Send mutated response body for second original message and trigger draining + new Thread( + () -> { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseBody( + BodyResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setBody( + ByteString.copyFromUtf8( + "mutated-msg-2")) + .build()) + .build()) + .build()) + .build()) + .setRequestDrain(true) + .build()); + try { + if (sidecarFinishLatch.await(5, TimeUnit.SECONDS)) { + responseObserver.onCompleted(); + } + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + }) + .start(); + } + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + drainCompletedLatch.countDown(); + } + }; + } + }; + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName) + .directExecutor() + .build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicReference> interceptedCallRef = + new AtomicReference<>(); + final List rawSentHeaders = new java.util.concurrent.CopyOnWriteArrayList<>(); + final List rawSentMessages = new java.util.concurrent.CopyOnWriteArrayList<>(); + final AtomicReference rawSentStatus = new AtomicReference<>(); + final AtomicReference rawSentTrailers = new AtomicReference<>(); + final CountDownLatch rawCloseLatch = new CountDownLatch(1); + + ServerCall rawCall = + new SimpleServerCall(METHOD_SAY_HELLO_RAW) { + @Override + public void sendHeaders(Metadata headers) { + rawSentHeaders.add(headers); + } + + @Override + public void sendMessage(InputStream message) { + rawSentMessages.add(message); + } + + @Override + public void close(Status status, Metadata trailers) { + rawSentStatus.set(status); + rawSentTrailers.set(trailers); + rawCloseLatch.countDown(); + } + + @Override + public boolean isReady() { + return true; + } + }; + + try { + interceptor.interceptCall( + rawCall, + new Metadata(), + (call, headers) -> { + interceptedCallRef.set(call); + return new ServerCall.Listener() {}; + }); + + ServerCall interceptedCall = interceptedCallRef.get(); + + // App sends response headers + interceptedCall.sendHeaders(new Metadata()); + // Wait for headers to be received by client + long startTime = System.currentTimeMillis(); + while (rawSentHeaders.isEmpty() && System.currentTimeMillis() - startTime < 5000) { + Thread.sleep(10); + } + assertThat(rawSentHeaders).hasSize(1); + + // App sends 1st message + interceptedCall.sendMessage( + new ByteArrayInputStream("original-msg-1".getBytes(StandardCharsets.UTF_8))); + // Wait for 1st mutated message to be received by client + startTime = System.currentTimeMillis(); + while (rawSentMessages.isEmpty() && System.currentTimeMillis() - startTime < 5000) { + Thread.sleep(10); + } + assertThat(rawSentMessages).hasSize(1); + + // App sends 2nd message + interceptedCall.sendMessage( + new ByteArrayInputStream("original-msg-2".getBytes(StandardCharsets.UTF_8))); + + // Wait for 2nd mutated message to be received by client, and wait for drain to be active + startTime = System.currentTimeMillis(); + while ((rawSentMessages.size() < 2 || interceptedCall.isReady()) + && System.currentTimeMillis() - startTime < 5000) { + Thread.sleep(10); + } + assertThat(rawSentMessages).hasSize(2); + assertThat(interceptedCall.isReady()).isFalse(); + + // Now that draining is active, App sends message and closes the server side of the call + // concurrently + final CountDownLatch appActionLatch = new CountDownLatch(1); + new Thread( + () -> { + interceptedCall.sendMessage( + new ByteArrayInputStream( + "unmutated-msg-during-drain".getBytes(StandardCharsets.UTF_8))); + interceptedCall.close(Status.OK, new Metadata()); + appActionLatch.countDown(); + }) + .start(); + + assertThat(appActionLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + // Assert that nothing has been delivered to client (rawCall) during active drain + assertThat(rawSentMessages).hasSize(2); // Still only 2 mutated messages + assertThat(rawSentStatus.get()).isNull(); + + // Now let sidecar complete the drain + sidecarFinishLatch.countDown(); + + // Wait for rawCall to close + assertThat(rawCloseLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + // Verify delivery order: headers, mutated-msg-1, mutated-msg-2, unmutated-msg-during-drain, + // and status OK + assertThat(rawSentHeaders).hasSize(1); + + List deliveredMessages = new ArrayList<>(); + for (InputStream is : rawSentMessages) { + deliveredMessages.add(new String(ByteStreams.toByteArray(is), StandardCharsets.UTF_8)); + } + assertThat(deliveredMessages) + .containsExactly("mutated-msg-1", "mutated-msg-2", "unmutated-msg-during-drain") + .inOrder(); + + assertThat(rawSentStatus.get().isOk()).isTrue(); + assertThat(rawSentTrailers.get()).isNotNull(); + } finally { + if (responseObserverRef.get() != null) { + try { + responseObserverRef.get().onCompleted(); + } catch (IllegalStateException ignored) { + // ignore + } + } + channelManager.close(); + } + } + + @Test + @SuppressWarnings("unchecked") + public void drainingStartsBeforeRequestHeaders_whenClientSendsMessages_thenBufferedAndDelivered() + throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + ExternalProcessor proto = + ExternalProcessor.newBuilder() + .setGrpcService( + GrpcService.newBuilder() + .setGoogleGrpc( + GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///" + uniqueExtProcServerName) + .addChannelCredentialsPlugin( + Any.newBuilder() + .setTypeUrl( + "type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build()) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestBodyMode(ProcessingMode.BodySendMode.GRPC) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + // External Processor Server + final CountDownLatch sidecarFinishLatch = new CountDownLatch(1); + final CountDownLatch drainCompletedLatch = new CountDownLatch(1); + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + responseObserverRef.set(responseObserver); + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + new Thread( + () -> { + responseObserver.onNext( + ProcessingResponse.newBuilder().setRequestDrain(true).build()); + try { + if (sidecarFinishLatch.await(5, TimeUnit.SECONDS)) { + responseObserver.onCompleted(); + } + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + }) + .start(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + drainCompletedLatch.countDown(); + } + }; + } + }; + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName) + .directExecutor() + .build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final List serverReceivedMessages = new java.util.concurrent.CopyOnWriteArrayList<>(); + final CountDownLatch appMessageLatch = new CountDownLatch(1); + ServerCall.Listener appListener = + new ServerCall.Listener() { + @Override + public void onMessage(InputStream message) { + try { + serverReceivedMessages.add( + new String(ByteStreams.toByteArray(message), StandardCharsets.UTF_8)); + appMessageLatch.countDown(); + } catch (IOException e) { + throw new RuntimeException(e); + } + } + }; + + final AtomicReference> interceptedCallRef = + new AtomicReference<>(); + ServerCall rawCall = + new SimpleServerCall(METHOD_SAY_HELLO_RAW) { + @Override + public boolean isReady() { + return true; + } + }; + + try { + ServerCall.Listener interceptedListener = + interceptor.interceptCall( + rawCall, + new Metadata(), + (call, headers) -> { + interceptedCallRef.set(call); + return appListener; + }); + + // Wait for drain to be processed + assertThat(drainCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(interceptedCallRef.get().isReady()).isFalse(); + + // Client sends request message during drain state + interceptedListener.onMessage( + new ByteArrayInputStream("client message during drain".getBytes(StandardCharsets.UTF_8))); + + // Verify app listener has NOT received the message yet because the drain is active + assertThat(serverReceivedMessages).isEmpty(); + + // Now let sidecar complete + sidecarFinishLatch.countDown(); + + // Wait for it to become ready again + long startTime = System.currentTimeMillis(); + while (!interceptedCallRef.get().isReady() && System.currentTimeMillis() - startTime < 5000) { + Thread.sleep(10); + } + assertThat(interceptedCallRef.get().isReady()).isTrue(); + + // Verify that the buffered client request message is now delivered to the app + assertThat(appMessageLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(serverReceivedMessages).containsExactly("client message during drain"); + } finally { + if (responseObserverRef.get() != null) { + try { + responseObserverRef.get().onCompleted(); + } catch (IllegalStateException ignored) { + // ignore + } + } + channelManager.close(); + } + } + + @Test + @SuppressWarnings("unchecked") + public void drainingStartsAfterRequestHeaders_whenClientSendsMessages_thenBufferedAndDelivered() + throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + ExternalProcessor proto = + ExternalProcessor.newBuilder() + .setGrpcService( + GrpcService.newBuilder() + .setGoogleGrpc( + GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///" + uniqueExtProcServerName) + .addChannelCredentialsPlugin( + Any.newBuilder() + .setTypeUrl( + "type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build()) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SEND) + .setRequestBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseBodyMode(ProcessingMode.BodySendMode.NONE) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + // External Processor Server + final CountDownLatch sidecarFinishLatch = new CountDownLatch(1); + final CountDownLatch drainCompletedLatch = new CountDownLatch(1); + final AtomicInteger extProcReceivedBodyCount = new AtomicInteger(0); + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + responseObserverRef.set(responseObserver); + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + } else if (request.hasRequestBody()) { + int bodyIdx = extProcReceivedBodyCount.incrementAndGet(); + if (bodyIdx == 1) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestBody( + BodyResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setBody( + ByteString.copyFromUtf8( + "mutated-msg-1")) + .build()) + .build()) + .build()) + .build()) + .build()); + } else if (bodyIdx == 2) { + new Thread( + () -> { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestBody( + BodyResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setBody( + ByteString.copyFromUtf8( + "mutated-msg-2")) + .build()) + .build()) + .build()) + .build()) + .setRequestDrain(true) + .build()); + try { + if (sidecarFinishLatch.await(5, TimeUnit.SECONDS)) { + responseObserver.onCompleted(); + } + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + }) + .start(); + } + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + drainCompletedLatch.countDown(); + } + }; + } + }; + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName) + .directExecutor() + .build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final List serverReceivedMessages = new java.util.concurrent.CopyOnWriteArrayList<>(); + final CountDownLatch appHalfCloseLatch = new CountDownLatch(1); + ServerCall.Listener appListener = + new ServerCall.Listener() { + @Override + public void onMessage(InputStream message) { + try { + serverReceivedMessages.add( + new String(ByteStreams.toByteArray(message), StandardCharsets.UTF_8)); + } catch (IOException e) { + throw new RuntimeException(e); + } + } + + @Override + public void onHalfClose() { + appHalfCloseLatch.countDown(); + } + }; + + final AtomicReference> interceptedCallRef = + new AtomicReference<>(); + ServerCall rawCall = + new SimpleServerCall(METHOD_SAY_HELLO_RAW) { + @Override + public boolean isReady() { + return true; + } + }; + + try { + ServerCall.Listener interceptedListener = + interceptor.interceptCall( + rawCall, + new Metadata(), + (call, headers) -> { + interceptedCallRef.set(call); + return appListener; + }); + + // Wait until the call is active (headers processed and delegate listener created) + long startTime = System.currentTimeMillis(); + while (interceptedCallRef.get() == null && System.currentTimeMillis() - startTime < 5000) { + Thread.sleep(10); + } + assertThat(interceptedCallRef.get()).isNotNull(); + + ServerCall interceptedCall = interceptedCallRef.get(); + + // Client sends 1st message + interceptedCall.request(1); + interceptedListener.onMessage( + new ByteArrayInputStream("original-msg-1".getBytes(StandardCharsets.UTF_8))); + // Wait for 1st mutated message to be received by app + startTime = System.currentTimeMillis(); + while (serverReceivedMessages.size() < 1 && System.currentTimeMillis() - startTime < 5000) { + Thread.sleep(10); + } + assertThat(serverReceivedMessages).containsExactly("mutated-msg-1"); + + // Client sends 2nd message + interceptedCall.request(1); + interceptedListener.onMessage( + new ByteArrayInputStream("original-msg-2".getBytes(StandardCharsets.UTF_8))); + // Wait for 2nd mutated message to be received by app, and wait for drain to be active + startTime = System.currentTimeMillis(); + while ((serverReceivedMessages.size() < 2 || interceptedCall.isReady()) + && System.currentTimeMillis() - startTime < 5000) { + Thread.sleep(10); + } + assertThat(serverReceivedMessages) + .containsExactly("mutated-msg-1", "mutated-msg-2") + .inOrder(); + assertThat(interceptedCall.isReady()).isFalse(); + + // Client concurrently sends message and half-closes during active drain + final CountDownLatch clientActionLatch = new CountDownLatch(1); + new Thread( + () -> { + interceptedListener.onMessage( + new ByteArrayInputStream( + "client-msg-during-drain".getBytes(StandardCharsets.UTF_8))); + interceptedListener.onHalfClose(); + clientActionLatch.countDown(); + }) + .start(); + + assertThat(clientActionLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + // Assert that nothing has been delivered to the app during active drain + assertThat(serverReceivedMessages).hasSize(2); // Still only 2 mutated messages + assertThat(appHalfCloseLatch.getCount()).isEqualTo(1); // Half-close not processed yet + + // Now let sidecar complete the drain + sidecarFinishLatch.countDown(); + + // Wait for the delegate half-close to be processed + assertThat(appHalfCloseLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + // Verify delivery order: mutated-msg-1, mutated-msg-2, client-msg-during-drain + assertThat(serverReceivedMessages) + .containsExactly("mutated-msg-1", "mutated-msg-2", "client-msg-during-drain") + .inOrder(); + } finally { + if (responseObserverRef.get() != null) { + try { + responseObserverRef.get().onCompleted(); + } catch (IllegalStateException ignored) { + // ignore + } + } + channelManager.close(); + } + } + + @Test + public void givenNoRequestDrain_whenExtProcStreamCompletesNormally_thenTreatsAsNonOkAndCallFails() + throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + ExternalProcessor proto = + createBaseProto(uniqueExtProcServerName) + .setFailureModeAllow(false) // Fail closed + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + // External Processor Server completes normally without sending request_drain = true + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + // Complete stream normally without drain + responseObserver.onCompleted(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + final AtomicReference closedStatus = new AtomicReference<>(); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + closedStatus.set(status); + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(closedStatus.get().getCode()).isEqualTo(Status.Code.INTERNAL); + assertThat(closedStatus.get().getDescription()).contains("External processor stream failed"); + channelManager.close(); + } + + // ============================================================================ + // Category 15: Inbound Backpressure (request(n) / pendingRequests) + // ============================================================================ + + @Test + @SuppressWarnings("unchecked") + public void givenObservabilityTrue_whenExtProcBusy_thenAppRequestsBuffered() throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + ExternalProcessor proto = + ExternalProcessor.newBuilder() + .setGrpcService( + GrpcService.newBuilder() + .setGoogleGrpc( + GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///" + uniqueExtProcServerName) + .addChannelCredentialsPlugin( + Any.newBuilder() + .setTypeUrl( + "type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build()) + .setObservabilityMode(true) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + // External Processor Server + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + responseObserverRef.set(responseObserver); + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + final AtomicBoolean sidecarReady = new AtomicBoolean(true); + final AtomicReference> sidecarListenerRef = + new AtomicReference<>(); + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName) + .directExecutor() + .intercept( + new io.grpc.ClientInterceptor() { + @Override + public io.grpc.ClientCall interceptCall( + MethodDescriptor method, + io.grpc.CallOptions callOptions, + io.grpc.Channel next) { + return new io.grpc.ForwardingClientCall.SimpleForwardingClientCall< + ReqT, RespT>(next.newCall(method, callOptions)) { + @Override + public void start( + Listener responseListener, Metadata headers) { + sidecarListenerRef.set( + (Listener) responseListener); + super.start(responseListener, headers); + } + + @Override + public boolean isReady() { + return sidecarReady.get(); + } + }; + } + }) + .build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicReference> interceptedCallRef = + new AtomicReference<>(); + ServerCallHandler nextHandler = + (call, headers) -> { + interceptedCallRef.set(call); + return new ServerCall.Listener() {}; + }; + + final AtomicInteger rawRequestCount = new AtomicInteger(0); + ServerCall rawCall = + new SimpleServerCall(METHOD_SAY_HELLO_RAW) { + @Override + public void request(int numMessages) { + rawRequestCount.addAndGet(numMessages); + } + }; + + try { + interceptor.interceptCall(rawCall, new Metadata(), nextHandler); + + // Wait for sidecar call to start and listener to be captured + long startTime = System.currentTimeMillis(); + while (sidecarListenerRef.get() == null && System.currentTimeMillis() - startTime < 5000) { + Thread.sleep(10); + } + assertThat(sidecarListenerRef.get()).isNotNull(); + + // Wait for activation + startTime = System.currentTimeMillis(); + while (!interceptedCallRef.get().isReady() && System.currentTimeMillis() - startTime < 5000) { + Thread.sleep(10); + } + assertThat(interceptedCallRef.get().isReady()).isTrue(); + + // Sidecar is busy + sidecarReady.set(false); + assertThat(interceptedCallRef.get().isReady()).isFalse(); + + // Application requests more messages + interceptedCallRef.get().request(5); + + // Verify raw call NOT requested yet + assertThat(rawRequestCount.get()).isEqualTo(0); + + // Sidecar becomes ready + sidecarReady.set(true); + sidecarListenerRef.get().onReady(); + + // After sidecar becomes ready, pending requests should be drained to raw call. + long start = System.currentTimeMillis(); + while (rawRequestCount.get() < 5 && System.currentTimeMillis() - start < 2000) { + Thread.sleep(10); + } + assertThat(rawRequestCount.get()).isEqualTo(5); + assertThat(interceptedCallRef.get().isReady()).isTrue(); + } finally { + if (responseObserverRef.get() != null) { + try { + responseObserverRef.get().onCompleted(); + } catch (IllegalStateException ignored) { + // Ignore if already closed + } + } + channelManager.close(); + } + } + + @Test + @SuppressWarnings("unchecked") + public void givenRequestDrainActive_whenAppRequestsMessages_thenRequestsBuffered() + throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + ExternalProcessor proto = + ExternalProcessor.newBuilder() + .setGrpcService( + GrpcService.newBuilder() + .setGoogleGrpc( + GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///" + uniqueExtProcServerName) + .addChannelCredentialsPlugin( + Any.newBuilder() + .setTypeUrl( + "type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch drainLatch = new CountDownLatch(1); + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + responseObserverRef.set(responseObserver); + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder().setRequestDrain(true).build()); + drainLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName) + .directExecutor() + .build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicReference> interceptedCallRef = + new AtomicReference<>(); + ServerCallHandler nextHandler = + (call, headers) -> { + interceptedCallRef.set(call); + return new ServerCall.Listener() {}; + }; + + final AtomicInteger rawRequestCount = new AtomicInteger(0); + ServerCall rawCall = + new SimpleServerCall(METHOD_SAY_HELLO_RAW) { + @Override + public void request(int numMessages) { + rawRequestCount.addAndGet(numMessages); + } + }; + + try { + interceptor.interceptCall(rawCall, new Metadata(), nextHandler); + + assertThat(drainLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + // Wait for interceptedCallRef to become ready first (it will transition to activated, then + // draining) + long start = System.currentTimeMillis(); + while (interceptedCallRef.get() == null && System.currentTimeMillis() - start < 5000) { + Thread.sleep(10); + } + assertThat(interceptedCallRef.get()).isNotNull(); + + // isReady() must return false during drain. + start = System.currentTimeMillis(); + while (interceptedCallRef.get().isReady() && System.currentTimeMillis() - start < 2000) { + Thread.sleep(10); + } + assertThat(interceptedCallRef.get().isReady()).isFalse(); + + // App requests more messages + interceptedCallRef.get().request(3); + + // isReady() should remain false during drain + assertThat(interceptedCallRef.get().isReady()).isFalse(); + + // Verify raw call NOT requested during drain + assertThat(rawRequestCount.get()).isEqualTo(0); + } finally { + if (responseObserverRef.get() != null) { + try { + responseObserverRef.get().onCompleted(); + } catch (IllegalStateException ignored) { + // Ignore if already closed + } + } + channelManager.close(); + } + } + + // ============================================================================ + // Category 16: Error Handling & Security + // ============================================================================ + + @Test + public void serverInterceptor_failOpen_allowsCallToProceed() throws Exception { + ExternalProcessor proto = createBaseProto(extProcServerName).setFailureModeAllow(true).build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + // External processor returns immediate error + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + responseObserver.onError( + Status.UNAVAILABLE.withDescription("ExtProc down").asRuntimeException()); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) {} + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicBoolean callStarted = new AtomicBoolean(false); + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + callStarted.set(true); + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callStarted.get()).isTrue(); + } + + @Test + public void serverInterceptor_failClosed_cancelsCall() throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setFailureModeAllow(false) // Fail closed + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + responseObserver.onError( + Status.INTERNAL.withDescription("Simulated sidecar failure").asRuntimeException()); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) {} + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + final AtomicReference closedStatus = new AtomicReference<>(); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + closedStatus.set(status); + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(closedStatus.get().getCode()).isEqualTo(Status.Code.INTERNAL); + assertThat(closedStatus.get().getDescription()).contains("External processor stream failed"); + } + + @Test + public void givenFailureModeAllowTrue_whenExtProcStreamFailsAfterRequestBodySent_thenCallFails() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setFailureModeAllow(true) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SEND) + .setRequestBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseBodyMode(ProcessingMode.BodySendMode.NONE) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SKIP) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = new CountDownLatch(1); + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + final StreamObserver responseObserver) { + responseObserverRef.set(responseObserver); + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + } else if (request.hasRequestBody()) { + // Fail the stream after receiving the request body + responseObserver.onError( + Status.INTERNAL + .withDescription("Simulated stream failure") + .asRuntimeException()); + extProcLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicBoolean callStarted = new AtomicBoolean(false); + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + callStarted.set(true); + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + final AtomicReference closedStatus = new AtomicReference<>(); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + closedStatus.set(status); + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + // Verify stream failed + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + // Verify call completed and failed with INTERNAL status (not fail-open) + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callStarted.get()).isFalse(); + assertThat(closedStatus.get().getCode()).isEqualTo(Status.Code.INTERNAL); + assertThat(closedStatus.get().getDescription()).contains("External processor stream failed"); + } + + @Test + public void givenFailureModeAllowTrue_whenExtProcStreamFailsAfterResponseBodySent_thenCallFails() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setFailureModeAllow(true) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setRequestBodyMode(ProcessingMode.BodySendMode.NONE) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SEND) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = new CountDownLatch(1); + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + final StreamObserver responseObserver) { + responseObserverRef.set(responseObserver); + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasResponseBody()) { + // Fail the stream after receiving the response body + responseObserver.onError( + Status.INTERNAL + .withDescription("Simulated stream failure") + .asRuntimeException()); + extProcLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicBoolean callStarted = new AtomicBoolean(false); + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + callStarted.set(true); + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + final AtomicReference closedStatus = new AtomicReference<>(); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + closedStatus.set(status); + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + // Verify call completed and failed with INTERNAL status (not fail-open) + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callStarted.get()).isTrue(); + assertThat(closedStatus.get().getCode()).isEqualTo(Status.Code.INTERNAL); + assertThat(closedStatus.get().getDescription()).contains("External processor stream failed"); + } + + @Test + public void givenObservabilityTrue_whenExtProcStreamFails_thenCallContinues() throws Exception { + ExternalProcessor proto = createBaseProto(extProcServerName).setObservabilityMode(true).build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + responseObserver.onError( + Status.UNAVAILABLE.withDescription("ExtProc down").asRuntimeException()); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) {} + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicBoolean callStarted = new AtomicBoolean(false); + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + callStarted.set(true); + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + final AtomicReference closedStatus = new AtomicReference<>(); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + closedStatus.set(status); + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + // Verify call completed and succeeded with OK status even though stream failed + assertThat(callStarted.get()).isTrue(); + assertThat(closedStatus.get().getCode()).isEqualTo(Status.Code.OK); + } + + // ============================================================================ + // Category 17: Immediate Response Handling + // ============================================================================ + + @Test + public void serverInterceptor_immediateResponse_whenReceived_thenDataPlaneCallClosed() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SEND) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SKIP) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setImmediateResponse( + io.envoyproxy.envoy.service.ext_proc.v3.ImmediateResponse.newBuilder() + .setGrpcStatus( + io.envoyproxy.envoy.service.ext_proc.v3.GrpcStatus + .newBuilder() + .setStatus(Status.UNAUTHENTICATED.getCode().value()) + .build()) + .setDetails("Custom security rejection") + .build()) + .build()); + responseObserver.onCompleted(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicBoolean dataPlaneStarted = new AtomicBoolean(false); + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + dataPlaneStarted.set(true); + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + final AtomicReference closedStatus = new AtomicReference<>(); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + closedStatus.set(status); + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(closedStatus.get().getCode()).isEqualTo(Status.Code.UNAUTHENTICATED); + assertThat(closedStatus.get().getDescription()).isEqualTo("Custom security rejection"); + assertThat(dataPlaneStarted.get()).isFalse(); + } + + @Test + public void serverInterceptor_immediateResponseDisabled_whenReceived_thenStreamErrored() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SEND) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SKIP) + .build()) + .setDisableImmediateResponse(true) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setImmediateResponse( + io.envoyproxy.envoy.service.ext_proc.v3.ImmediateResponse.newBuilder() + .setGrpcStatus( + io.envoyproxy.envoy.service.ext_proc.v3.GrpcStatus + .newBuilder() + .setStatus(Status.UNAUTHENTICATED.getCode().value()) + .build()) + .build()) + .build()); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + final AtomicReference closedStatus = new AtomicReference<>(); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + closedStatus.set(status); + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(closedStatus.get().getCode()).isAnyOf(Status.Code.INTERNAL, Status.Code.UNAVAILABLE); + } + + @Test + @SuppressWarnings("FutureReturnValueIgnored") + public void + serverInterceptor_pendingData_whenImmediateResponseReceived_thenDeliversDataBeforeStatus() + throws Exception { + final String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + final List clientEvents = Collections.synchronizedList(new ArrayList<>()); + final CountDownLatch finishLatch = new CountDownLatch(1); + final CountDownLatch extProcCompletedLatch = new CountDownLatch(1); + final java.util.concurrent.ExecutorService extProcResponseExecutor = + java.util.concurrent.Executors.newSingleThreadExecutor(); + final Metadata.Key immediateKey = + Metadata.Key.of("x-immediate-header", Metadata.ASCII_STRING_MARSHALLER); + final AtomicReference clientTrailers = new AtomicReference<>(); + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + final StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + extProcResponseExecutor.submit( + () -> { + synchronized (responseObserver) { + if (request.hasRequestBody()) { + try { + Thread.sleep(500); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setImmediateResponse( + io.envoyproxy.envoy.service.ext_proc.v3.ImmediateResponse + .newBuilder() + .setGrpcStatus( + io.envoyproxy.envoy.service.ext_proc.v3.GrpcStatus + .newBuilder() + .setStatus( + Status.UNAUTHENTICATED.getCode().value()) + .build()) + .setDetails("Immediate Auth Failure") + .setHeaders( + io.envoyproxy.envoy.service.ext_proc.v3.HeaderMutation + .newBuilder() + .addSetHeaders( + io.envoyproxy.envoy.config.core.v3 + .HeaderValueOption.newBuilder() + .setHeader( + io.envoyproxy.envoy.config.core.v3 + .HeaderValue.newBuilder() + .setKey("x-immediate-header") + .setValue("true") + .build()) + .build()) + .build()) + .build()) + .build()); + responseObserver.onCompleted(); + } + } + }); + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + extProcCompletedLatch.countDown(); + } + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName) + .directExecutor() + .build())); + + ExternalProcessor proto = + ExternalProcessor.newBuilder() + .setGrpcService( + GrpcService.newBuilder() + .setGoogleGrpc( + GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///" + uniqueExtProcServerName) + .addChannelCredentialsPlugin( + Any.newBuilder() + .setTypeUrl( + "type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build()) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setRequestBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseBodyMode(ProcessingMode.BodySendMode.NONE) + .build()) + .build(); + + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public StreamObserver sayHelloBidi( + StreamObserver responseObserver) { + byte[] messageBytes = "server-response".getBytes(StandardCharsets.UTF_8); + responseObserver.onNext(new ByteArrayInputStream(messageBytes)); + responseObserver.onCompleted(); + + return new StreamObserver() { + @Override + public void onNext(InputStream value) {} + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_BIDI, io.grpc.CallOptions.DEFAULT); + + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onHeaders(Metadata headers) { + clientEvents.add("HEADERS"); + } + + @Override + public void onMessage(InputStream message) { + clientEvents.add("MESSAGE"); + } + + @Override + public void onClose(Status status, Metadata trailers) { + clientEvents.add("CLOSE:" + status.getCode()); + clientTrailers.set(trailers); + finishLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + byte[] requestBytes = "request-body".getBytes(StandardCharsets.UTF_8); + clientCall.sendMessage(new ByteArrayInputStream(requestBytes)); + clientCall.halfClose(); + + assertThat(finishLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(clientEvents).containsExactly("HEADERS", "MESSAGE", "CLOSE:UNAUTHENTICATED"); + assertThat(clientTrailers.get().get(immediateKey)).isEqualTo("true"); + + extProcResponseExecutor.shutdown(); + channelManager.close(); + } + + // ============================================================================ + // Category 18: Resource Management + // ============================================================================ + + @Test + public void givenFilter_whenClosed_thenCachedChannelManagerIsClosed() throws Exception { + CachedChannelManager mockChannelManager = Mockito.mock(CachedChannelManager.class); + ExternalProcessorFilter filter = new ExternalProcessorFilter(FAKE_CONTEXT, mockChannelManager); + filter.close(); + Mockito.verify(mockChannelManager).close(); + } + + // ============================================================================ + // Category 19: Data plane rpc cancellation + // ============================================================================ + + @Test + @SuppressWarnings("unchecked") + public void givenActiveRpc_whenDataPlaneCallCancelled_thenExtProcStreamIsErrored() + throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + ExternalProcessor proto = + ExternalProcessor.newBuilder() + .setGrpcService( + GrpcService.newBuilder() + .setGoogleGrpc( + GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///" + uniqueExtProcServerName) + .addChannelCredentialsPlugin( + Any.newBuilder() + .setTypeUrl( + "type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch cancelLatch = new CountDownLatch(1); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) {} + + @Override + public void onError(Throwable t) { + cancelLatch.countDown(); + } + + @Override + public void onCompleted() {} + }; + } + }; + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName) + .directExecutor() + .build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicReference> interceptedListenerRef = + new AtomicReference<>(); + ServerCallHandler nextHandler = + (call, headers) -> { + ServerCall.Listener listener = new ServerCall.Listener() {}; + interceptedListenerRef.set(listener); + return listener; + }; + + ServerCall rawCall = new SimpleServerCall(METHOD_SAY_HELLO_RAW); + + try { + ServerCall.Listener listener = + interceptor.interceptCall(rawCall, new Metadata(), nextHandler); + + // Client cancels the RPC + listener.onCancel(); + + // Verify sidecar control stream also received onError/cancellation + assertThat(cancelLatch.await(5, TimeUnit.SECONDS)).isTrue(); + } finally { + channelManager.close(); + } + } + + // ============================================================================ + // Category 20: Flow Control when side stream is full + // ============================================================================ + + @Test + @SuppressWarnings("unchecked") + public void givenObservabilityModeFalse_whenExtProcBusy_thenIsReadyReturnsFalse() + throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + ExternalProcessor proto = + ExternalProcessor.newBuilder() + .setGrpcService( + GrpcService.newBuilder() + .setGoogleGrpc( + GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///" + uniqueExtProcServerName) + .addChannelCredentialsPlugin( + Any.newBuilder() + .setTypeUrl( + "type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build()) + .setObservabilityMode(false) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + // External Processor Server + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + responseObserverRef.set(responseObserver); + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + final AtomicBoolean sidecarReady = new AtomicBoolean(true); + final AtomicReference> sidecarListenerRef = + new AtomicReference<>(); + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName) + .directExecutor() + .intercept( + new io.grpc.ClientInterceptor() { + @Override + public io.grpc.ClientCall interceptCall( + MethodDescriptor method, + io.grpc.CallOptions callOptions, + io.grpc.Channel next) { + return new io.grpc.ForwardingClientCall.SimpleForwardingClientCall< + ReqT, RespT>(next.newCall(method, callOptions)) { + @Override + public void start( + Listener responseListener, Metadata headers) { + sidecarListenerRef.set( + (Listener) responseListener); + super.start(responseListener, headers); + } + + @Override + public boolean isReady() { + return sidecarReady.get(); + } + }; + } + }) + .build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicReference> interceptedCallRef = + new AtomicReference<>(); + ServerCallHandler nextHandler = + (call, headers) -> { + interceptedCallRef.set(call); + return new ServerCall.Listener() {}; + }; + + final AtomicBoolean rawCallReady = new AtomicBoolean(true); + ServerCall rawCall = + new SimpleServerCall(METHOD_SAY_HELLO_RAW) { + @Override + public boolean isReady() { + return rawCallReady.get(); + } + }; + + try { + interceptor.interceptCall(rawCall, new Metadata(), nextHandler); + + // Wait for sidecar call to start and listener to be captured + long startTime = System.currentTimeMillis(); + while (sidecarListenerRef.get() == null && System.currentTimeMillis() - startTime < 5000) { + Thread.sleep(10); + } + assertThat(sidecarListenerRef.get()).isNotNull(); + + // Wait for activation + startTime = System.currentTimeMillis(); + while (!interceptedCallRef.get().isReady() && System.currentTimeMillis() - startTime < 5000) { + Thread.sleep(10); + } + assertThat(interceptedCallRef.get().isReady()).isTrue(); + + // Sidecar is busy -> intercepted call becomes busy + sidecarReady.set(false); + assertThat(interceptedCallRef.get().isReady()).isFalse(); + + // Sidecar becomes ready, raw call is busy -> intercepted call is STILL ready (observability + // mode false) + sidecarReady.set(true); + rawCallReady.set(false); + assertThat(interceptedCallRef.get().isReady()).isTrue(); + } finally { + if (responseObserverRef.get() != null) { + try { + responseObserverRef.get().onCompleted(); + } catch (IllegalStateException ignored) { + // Ignore if already closed + } + } + channelManager.close(); + } + } + + @Test + @SuppressWarnings("unchecked") + public void givenObservabilityModeFalse_whenExtProcBusy_thenAppRequestsAreBuffered() + throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + ExternalProcessor proto = + ExternalProcessor.newBuilder() + .setGrpcService( + GrpcService.newBuilder() + .setGoogleGrpc( + GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///" + uniqueExtProcServerName) + .addChannelCredentialsPlugin( + Any.newBuilder() + .setTypeUrl( + "type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build()) + .setObservabilityMode(false) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + // External Processor Server + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + final StreamObserver responseObserver) { + responseObserverRef.set(responseObserver); + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + grpcCleanup.register( + InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + final AtomicBoolean sidecarReady = new AtomicBoolean(true); + final AtomicReference> sidecarListenerRef = + new AtomicReference<>(); + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName) + .directExecutor() + .intercept( + new io.grpc.ClientInterceptor() { + @Override + public io.grpc.ClientCall interceptCall( + MethodDescriptor method, + io.grpc.CallOptions callOptions, + io.grpc.Channel next) { + return new io.grpc.ForwardingClientCall.SimpleForwardingClientCall< + ReqT, RespT>(next.newCall(method, callOptions)) { + @Override + public void start( + Listener responseListener, Metadata headers) { + sidecarListenerRef.set( + (Listener) responseListener); + super.start(responseListener, headers); + } + + @Override + public boolean isReady() { + return sidecarReady.get(); + } + }; + } + }) + .build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicReference> interceptedCallRef = + new AtomicReference<>(); + ServerCallHandler nextHandler = + (call, headers) -> { + interceptedCallRef.set(call); + return new ServerCall.Listener() {}; + }; + + final AtomicInteger rawRequestCount = new AtomicInteger(0); + ServerCall rawCall = + new SimpleServerCall(METHOD_SAY_HELLO_RAW) { + @Override + public void request(int numMessages) { + rawRequestCount.addAndGet(numMessages); + } + }; + + try { + interceptor.interceptCall(rawCall, new Metadata(), nextHandler); + + // Wait for sidecar call to start and listener to be captured + long startTime = System.currentTimeMillis(); + while (sidecarListenerRef.get() == null && System.currentTimeMillis() - startTime < 5000) { + Thread.sleep(10); + } + assertThat(sidecarListenerRef.get()).isNotNull(); + + // Wait for activation + startTime = System.currentTimeMillis(); + while (!interceptedCallRef.get().isReady() && System.currentTimeMillis() - startTime < 5000) { + Thread.sleep(10); + } + assertThat(interceptedCallRef.get().isReady()).isTrue(); + + // Sidecar becomes busy -> request(5) should be buffered + sidecarReady.set(false); + interceptedCallRef.get().request(5); + assertThat(rawRequestCount.get()).isEqualTo(0); + + // Sidecar becomes ready -> buffered requests should be drained + sidecarReady.set(true); + sidecarListenerRef.get().onReady(); + + long start = System.currentTimeMillis(); + while (rawRequestCount.get() < 5 && System.currentTimeMillis() - start < 2000) { + Thread.sleep(10); + } + assertThat(rawRequestCount.get()).isEqualTo(5); + } finally { + if (responseObserverRef.get() != null) { + try { + responseObserverRef.get().onCompleted(); + } catch (IllegalStateException ignored) { + // Ignore if already closed + } + } + channelManager.close(); + } + } + + // ============================================================================ + // Category 21: Streaming Completeness (Client & Bi-Di) + // ============================================================================ + + @Test + public void serverInterceptor_clientStreaming_streamCompletenessSucceeds() throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SEND) + .setRequestBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SEND) + .setResponseBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SEND) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final List capturedRequests = + java.util.Collections.synchronizedList(new ArrayList<>()); + final CountDownLatch sidecarFinishedLatch = new CountDownLatch(1); + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + capturedRequests.add(request); + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + } else if (request.hasRequestBody()) { + BodyResponse.Builder bodyResponse = BodyResponse.newBuilder(); + if (request.getRequestBody().getBody().isEmpty() + && (request.getRequestBody().getEndOfStream() + || request.getRequestBody().getEndOfStreamWithoutMessage())) { + bodyResponse.setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setEndOfStream( + request.getRequestBody().getEndOfStream()) + .setEndOfStreamWithoutMessage( + request + .getRequestBody() + .getEndOfStreamWithoutMessage()) + .build()) + .build()) + .build()); + } else { + bodyResponse.setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setBody(request.getRequestBody().getBody()) + .setEndOfStream( + request.getRequestBody().getEndOfStream()) + .build()) + .build()) + .build()); + } + responseObserver.onNext( + ProcessingResponse.newBuilder().setRequestBody(bodyResponse).build()); + } else if (request.hasResponseHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders(HeadersResponse.newBuilder().build()) + .build()); + } else if (request.hasResponseBody()) { + BodyResponse.Builder bodyResponse = BodyResponse.newBuilder(); + if (request.getResponseBody().getBody().isEmpty() + && (request.getResponseBody().getEndOfStream() + || request.getResponseBody().getEndOfStreamWithoutMessage())) { + bodyResponse.setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse(StreamedBodyResponse.newBuilder().build()) + .build()) + .build()); + } else { + bodyResponse.setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setBody(request.getResponseBody().getBody()) + .build()) + .build()) + .build()); + } + responseObserver.onNext( + ProcessingResponse.newBuilder().setResponseBody(bodyResponse).build()); + } else if (request.hasResponseTrailers()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseTrailers(TrailersResponse.newBuilder().build()) + .build()); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + sidecarFinishedLatch.countDown(); + } + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final List receivedDataPlaneRequests = + java.util.Collections.synchronizedList(new ArrayList<>()); + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public StreamObserver sayHelloClientStreaming( + StreamObserver responseObserver) { + return new StreamObserver() { + @Override + public void onNext(InputStream value) { + try { + byte[] bytes = ByteStreams.toByteArray(value); + receivedDataPlaneRequests.add(new String(bytes, StandardCharsets.UTF_8)); + } catch (IOException e) { + responseObserver.onError(e); + } + } + + @Override + public void onError(Throwable t) { + responseObserver.onError(t); + } + + @Override + public void onCompleted() { + responseObserver.onNext( + new ByteArrayInputStream("response-payload".getBytes(StandardCharsets.UTF_8))); + responseObserver.onCompleted(); + } + }; + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_CLIENT_STREAMING, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch clientCompletedLatch = new CountDownLatch(1); + final AtomicReference closedStatus = new AtomicReference<>(); + final List clientResponses = java.util.Collections.synchronizedList(new ArrayList<>()); + + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onMessage(InputStream message) { + try { + byte[] bytes = ByteStreams.toByteArray(message); + clientResponses.add(new String(bytes, StandardCharsets.UTF_8)); + } catch (IOException e) { + throw new RuntimeException(e); + } + } + + @Override + public void onClose(Status status, Metadata trailers) { + closedStatus.set(status); + clientCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(100); + clientCall.sendMessage(new ByteArrayInputStream("msg-1".getBytes(StandardCharsets.UTF_8))); + clientCall.sendMessage(new ByteArrayInputStream("msg-2".getBytes(StandardCharsets.UTF_8))); + clientCall.sendMessage(new ByteArrayInputStream("msg-3".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(clientCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(closedStatus.get().isOk()).isTrue(); + + assertThat(receivedDataPlaneRequests).containsExactly("msg-1", "msg-2", "msg-3").inOrder(); + assertThat(clientResponses).containsExactly("response-payload"); + + assertThat(sidecarFinishedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + int requestHeadersCount = 0; + int requestBodyCount = 0; + for (ProcessingRequest request : capturedRequests) { + if (request.hasRequestHeaders()) { + requestHeadersCount++; + } else if (request.hasRequestBody()) { + requestBodyCount++; + } + } + assertThat(requestHeadersCount).isEqualTo(1); + assertThat(requestBodyCount).isEqualTo(4); + } + + // ============================================================================ + // Category 22: Header Forwarding Rules + // ============================================================================ + + @Test + public void serverInterceptor_allowedHeaders_whenHeadersForwarded_thenOnlyAllowedAreSent() + throws Exception { + final AtomicReference capturedHeaders = + new AtomicReference<>(); + final CountDownLatch sidecarLatch = new CountDownLatch(1); + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + capturedHeaders.set(request.getRequestHeaders()); + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + sidecarLatch.countDown(); + } else if (request.hasResponseHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders(HeadersResponse.newBuilder().build()) + .build()); + } else if (request.hasRequestBody()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestBody( + io.envoyproxy.envoy.service.ext_proc.v3.BodyResponse.newBuilder() + .build()) + .build()); + } else if (request.hasResponseBody()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseBody( + io.envoyproxy.envoy.service.ext_proc.v3.BodyResponse.newBuilder() + .build()) + .build()); + } else if (request.hasRequestTrailers()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestTrailers( + io.envoyproxy.envoy.service.ext_proc.v3.TrailersResponse.newBuilder() + .build()) + .build()); + } else if (request.hasResponseTrailers()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseTrailers( + io.envoyproxy.envoy.service.ext_proc.v3.TrailersResponse.newBuilder() + .build()) + .build()); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + // Config with forward_rules and explicit processing mode: + // allowed_headers = ["x-allowed-*", "content-type"], requestHeaderMode = SEND, others = SKIP + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SEND) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SKIP) + .build()) + .setForwardRules( + HeaderForwardingRules.newBuilder() + .setAllowedHeaders( + io.envoyproxy.envoy.type.matcher.v3.ListStringMatcher.newBuilder() + .addPatterns( + io.envoyproxy.envoy.type.matcher.v3.StringMatcher.newBuilder() + .setPrefix("x-allowed-") + .build()) + .addPatterns( + io.envoyproxy.envoy.type.matcher.v3.StringMatcher.newBuilder() + .setExact("content-type") + .build()) + .build()) + .build()) + .build(); + + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + final AtomicReference callStatus = new AtomicReference<>(); + Metadata headers = new Metadata(); + headers.put(Metadata.Key.of("x-allowed-1", Metadata.ASCII_STRING_MARSHALLER), "v1"); + headers.put(Metadata.Key.of("x-disallowed", Metadata.ASCII_STRING_MARSHALLER), "v2"); + headers.put( + Metadata.Key.of("content-type", Metadata.ASCII_STRING_MARSHALLER), "application/grpc"); + + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callStatus.set(status); + callCompletedLatch.countDown(); + } + }, + headers); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + boolean sidecarAwaited = sidecarLatch.await(5, TimeUnit.SECONDS); + boolean completedAwaited = callCompletedLatch.await(5, TimeUnit.SECONDS); + + assertThat(sidecarAwaited).isTrue(); + assertThat(completedAwaited).isTrue(); + assertThat(callStatus.get().isOk()).isTrue(); + + List headerKeys = new ArrayList<>(); + for (io.envoyproxy.envoy.config.core.v3.HeaderValue hv : + capturedHeaders.get().getHeaders().getHeadersList()) { + headerKeys.add(hv.getKey()); + } + + assertThat(headerKeys).contains("x-allowed-1"); + assertThat(headerKeys).contains("content-type"); + assertThat(headerKeys).doesNotContain("x-disallowed"); + } + + @Test + public void serverInterceptor_disallowedHeaders_whenHeadersForwarded_thenSkipped() + throws Exception { + final AtomicReference capturedHeaders = + new AtomicReference<>(); + final CountDownLatch sidecarLatch = new CountDownLatch(1); + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + capturedHeaders.set(request.getRequestHeaders()); + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + sidecarLatch.countDown(); + } else if (request.hasResponseHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders(HeadersResponse.newBuilder().build()) + .build()); + } else if (request.hasRequestBody()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestBody( + io.envoyproxy.envoy.service.ext_proc.v3.BodyResponse.newBuilder() + .build()) + .build()); + } else if (request.hasResponseBody()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseBody( + io.envoyproxy.envoy.service.ext_proc.v3.BodyResponse.newBuilder() + .build()) + .build()); + } else if (request.hasRequestTrailers()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestTrailers( + io.envoyproxy.envoy.service.ext_proc.v3.TrailersResponse.newBuilder() + .build()) + .build()); + } else if (request.hasResponseTrailers()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseTrailers( + io.envoyproxy.envoy.service.ext_proc.v3.TrailersResponse.newBuilder() + .build()) + .build()); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + // Config with forward_rules and explicit processing mode: requestHeaderMode = SEND, others = + // SKIP + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SEND) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SKIP) + .build()) + .setForwardRules( + HeaderForwardingRules.newBuilder() + .setDisallowedHeaders( + io.envoyproxy.envoy.type.matcher.v3.ListStringMatcher.newBuilder() + .addPatterns( + io.envoyproxy.envoy.type.matcher.v3.StringMatcher.newBuilder() + .setExact("x-secret") + .build()) + .addPatterns( + io.envoyproxy.envoy.type.matcher.v3.StringMatcher.newBuilder() + .setExact("authorization") + .build()) + .build()) + .build()) + .build(); + + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + Metadata headers = new Metadata(); + headers.put(Metadata.Key.of("x-foo", Metadata.ASCII_STRING_MARSHALLER), "v1"); + headers.put(Metadata.Key.of("x-secret", Metadata.ASCII_STRING_MARSHALLER), "v2"); + headers.put(Metadata.Key.of("authorization", Metadata.ASCII_STRING_MARSHALLER), "v3"); + + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + headers); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(sidecarLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + List headerKeys = new ArrayList<>(); + for (io.envoyproxy.envoy.config.core.v3.HeaderValue hv : + capturedHeaders.get().getHeaders().getHeadersList()) { + headerKeys.add(hv.getKey()); + } + + assertThat(headerKeys).contains("x-foo"); + assertThat(headerKeys).doesNotContain("x-secret"); + assertThat(headerKeys).doesNotContain("authorization"); + } + + @Test + public void serverInterceptor_bothRules_whenHeadersForwarded_thenBothAreApplied() + throws Exception { + final AtomicReference capturedHeaders = + new AtomicReference<>(); + final CountDownLatch sidecarLatch = new CountDownLatch(1); + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + capturedHeaders.set(request.getRequestHeaders()); + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + sidecarLatch.countDown(); + } else if (request.hasResponseHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders(HeadersResponse.newBuilder().build()) + .build()); + } else if (request.hasRequestBody()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestBody( + io.envoyproxy.envoy.service.ext_proc.v3.BodyResponse.newBuilder() + .build()) + .build()); + } else if (request.hasResponseBody()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseBody( + io.envoyproxy.envoy.service.ext_proc.v3.BodyResponse.newBuilder() + .build()) + .build()); + } else if (request.hasRequestTrailers()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestTrailers( + io.envoyproxy.envoy.service.ext_proc.v3.TrailersResponse.newBuilder() + .build()) + .build()); + } else if (request.hasResponseTrailers()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseTrailers( + io.envoyproxy.envoy.service.ext_proc.v3.TrailersResponse.newBuilder() + .build()) + .build()); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + // Config with forward_rules and explicit processing mode: requestHeaderMode = SEND, others = + // SKIP + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SEND) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SKIP) + .build()) + .setForwardRules( + HeaderForwardingRules.newBuilder() + .setAllowedHeaders( + io.envoyproxy.envoy.type.matcher.v3.ListStringMatcher.newBuilder() + .addPatterns( + io.envoyproxy.envoy.type.matcher.v3.StringMatcher.newBuilder() + .setPrefix("x-foo-") + .build()) + .build()) + .setDisallowedHeaders( + io.envoyproxy.envoy.type.matcher.v3.ListStringMatcher.newBuilder() + .addPatterns( + io.envoyproxy.envoy.type.matcher.v3.StringMatcher.newBuilder() + .setExact("x-foo-secret") + .build()) + .build()) + .build()) + .build(); + + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + Metadata headers = new Metadata(); + headers.put(Metadata.Key.of("x-foo-1", Metadata.ASCII_STRING_MARSHALLER), "v1"); + headers.put(Metadata.Key.of("x-foo-secret", Metadata.ASCII_STRING_MARSHALLER), "v2"); + headers.put(Metadata.Key.of("x-bar", Metadata.ASCII_STRING_MARSHALLER), "v3"); + + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + headers); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(sidecarLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + List headerKeys = new ArrayList<>(); + for (io.envoyproxy.envoy.config.core.v3.HeaderValue hv : + capturedHeaders.get().getHeaders().getHeadersList()) { + headerKeys.add(hv.getKey()); + } + + assertThat(headerKeys).contains("x-foo-1"); + assertThat(headerKeys).doesNotContain("x-foo-secret"); + assertThat(headerKeys).doesNotContain("x-bar"); + } + + // ============================================================================ + // Category 23: Response Ordering Checks + // ============================================================================ + + @Test + public void serverInterceptor_outOfOrderResponses_whenMessageArrivesBeforeHeaders_thenFails() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SEND) + .setResponseBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SEND) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch sidecarLatch = new CountDownLatch(1); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders( + HeadersResponse.newBuilder() + .setResponse(CommonResponse.newBuilder().build()) + .build()) + .build()); + } else if (request.hasResponseHeaders()) { + // Violate order: send ResponseBody response when ResponseHeaders response is + // expected + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseBody( + BodyResponse.newBuilder() + .setResponse(CommonResponse.newBuilder().build()) + .build()) + .build()); + sidecarLatch.countDown(); + responseObserver.onCompleted(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + final AtomicReference receivedStatus = new AtomicReference<>(); + + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + receivedStatus.set(status); + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(sidecarLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(receivedStatus.get().getCode()).isEqualTo(Status.Code.INTERNAL); + assertThat(receivedStatus.get().getDescription()).contains("External processor stream failed"); + } + + @Test + public void serverInterceptor_validOrder_whenResponsesArriveInOrder_thenSucceeds() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SEND) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = + new CountDownLatch(2); // 1 for request headers, 1 for response headers + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders( + HeadersResponse.newBuilder() + .setResponse(CommonResponse.newBuilder().build()) + .build()) + .build()); + extProcLatch.countDown(); + } else if (request.hasResponseHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders( + HeadersResponse.newBuilder() + .setResponse(CommonResponse.newBuilder().build()) + .build()) + .build()); + responseObserver.onCompleted(); + extProcLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + final AtomicReference receivedStatus = new AtomicReference<>(); + + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + receivedStatus.set(status); + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(receivedStatus.get().getCode()).isEqualTo(Status.Code.OK); + } + + // ============================================================================ + // Category 24: Header Response Status Checks + // ============================================================================ + + @Test + public void serverInterceptor_requestHeadersResponse_whenStatusIsContinueAndReplace_thenFails() + throws Exception { + ExternalProcessor proto = createBaseProto(extProcServerName).build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch sidecarLatch = new CountDownLatch(1); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders( + HeadersResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setStatus( + CommonResponse.ResponseStatus.CONTINUE_AND_REPLACE) + .build()) + .build()) + .build()); + sidecarLatch.countDown(); + responseObserver.onCompleted(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + final AtomicReference receivedStatus = new AtomicReference<>(); + + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + receivedStatus.set(status); + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(sidecarLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(receivedStatus.get().getCode()).isEqualTo(Status.Code.INTERNAL); + assertThat(receivedStatus.get().getDescription()).contains("External processor stream failed"); + } + + @Test + public void serverInterceptor_responseHeadersResponse_whenStatusIsContinueAndReplace_thenFails() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SEND) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch sidecarLatch = new CountDownLatch(1); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + } else if (request.hasResponseHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders( + HeadersResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setStatus( + CommonResponse.ResponseStatus.CONTINUE_AND_REPLACE) + .build()) + .build()) + .build()); + sidecarLatch.countDown(); + responseObserver.onCompleted(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext(request); + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + final AtomicReference receivedStatus = new AtomicReference<>(); + + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + receivedStatus.set(status); + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(sidecarLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(receivedStatus.get().getCode()).isEqualTo(Status.Code.INTERNAL); + assertThat(receivedStatus.get().getDescription()).contains("External processor stream failed"); + } + + // ============================================================================ + // Category 25: Concurrency and Thread Safety (Serialization) + // ============================================================================ + + private static class ConcurrencyDetectingServerCall + extends io.grpc.ForwardingServerCall.SimpleForwardingServerCall { + private final AtomicInteger activeCalls = new AtomicInteger(0); + private final AtomicBoolean concurrentCallDetected = new AtomicBoolean(false); + + ConcurrencyDetectingServerCall(ServerCall delegate) { + super(delegate); + } + + @Override + public void sendMessage(InputStream message) { + int active = activeCalls.incrementAndGet(); + if (active > 1) { + concurrentCallDetected.set(true); + } + try { + Thread.sleep(50); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + super.sendMessage(message); + activeCalls.decrementAndGet(); + } + + @Override + public void sendHeaders(Metadata headers) { + int active = activeCalls.incrementAndGet(); + if (active > 1) { + concurrentCallDetected.set(true); + } + try { + Thread.sleep(50); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + super.sendHeaders(headers); + activeCalls.decrementAndGet(); + } + } + + @Test + public void serverInterceptor_concurrency_serializesDelegateCallbacks() throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestBodyMode(ProcessingMode.BodySendMode.GRPC) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final int iterations = 500; + final CountDownLatch clientDone = new CountDownLatch(1); + final AtomicBoolean raceDetected = new AtomicBoolean(false); + final AtomicInteger activeServiceCalls = new AtomicInteger(0); + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders( + HeadersResponse.newBuilder() + .setResponse(CommonResponse.newBuilder().build()) + .build()) + .build()); + } else if (request.hasRequestBody()) { + boolean eos = + request.getRequestBody().getEndOfStream() + || request.getRequestBody().getEndOfStreamWithoutMessage(); + BodyResponse.Builder bodyResponseBuilder = BodyResponse.newBuilder(); + if (eos) { + bodyResponseBuilder.setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setEndOfStream( + request.getRequestBody().getEndOfStream()) + .setEndOfStreamWithoutMessage( + request + .getRequestBody() + .getEndOfStreamWithoutMessage()) + .build()) + .build()) + .build()); + } else { + ByteString body = request.getRequestBody().getBody(); + bodyResponseBuilder.setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder().setBody(body).build()) + .build()) + .build()); + } + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestBody(bodyResponseBuilder.build()) + .build()); + } else if (request.hasResponseHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders( + HeadersResponse.newBuilder() + .setResponse(CommonResponse.newBuilder().build()) + .build()) + .build()); + } else if (request.hasResponseBody()) { + boolean eos = + request.getResponseBody().getEndOfStream() + || request.getResponseBody().getEndOfStreamWithoutMessage(); + BodyResponse.Builder bodyResponseBuilder = BodyResponse.newBuilder(); + if (eos) { + bodyResponseBuilder.setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse(StreamedBodyResponse.newBuilder().build()) + .build()) + .build()); + } else { + ByteString body = request.getResponseBody().getBody(); + bodyResponseBuilder.setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder().setBody(body).build()) + .build()) + .build()); + } + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseBody(bodyResponseBuilder.build()) + .build()); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public StreamObserver sayHelloBidi( + StreamObserver responseObserver) { + return new StreamObserver() { + @Override + public void onNext(InputStream value) { + int active = activeServiceCalls.incrementAndGet(); + if (active > 1) { + raceDetected.set(true); + } + try { + Thread.sleep(1); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + activeServiceCalls.decrementAndGet(); + responseObserver.onNext(value); + } + + @Override + public void onError(Throwable t) { + responseObserver.onError(t); + } + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_BIDI, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch clientClosedLatch = new CountDownLatch(1); + final AtomicInteger clientReceivedMessages = new AtomicInteger(0); + + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onMessage(InputStream message) { + clientReceivedMessages.incrementAndGet(); + clientCall.request(1); + } + + @Override + public void onClose(Status status, Metadata trailers) { + clientClosedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + + Thread clientThread = + new Thread( + () -> { + try { + for (int i = 0; i < iterations; i++) { + clientCall.sendMessage(new ByteArrayInputStream(new byte[10])); + Thread.sleep(0, 10000); + } + clientCall.halfClose(); + clientDone.countDown(); + } catch (Exception e) { + e.printStackTrace(); + raceDetected.set(true); + } + }); + + clientThread.start(); + clientThread.join(); + + assertThat(clientDone.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(clientClosedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(clientReceivedMessages.get()).isEqualTo(iterations); + assertThat(raceDetected.get()).isFalse(); + } + + @Test + public void serverInterceptor_outboundStreamTermination_serializesSendMessage() throws Exception { + // Configure response body mode to GRPC + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setFailureModeAllow(true) + .setProcessingMode( + ProcessingMode.newBuilder() + .setResponseBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SEND) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + final CountDownLatch streamActiveLatch = new CountDownLatch(1); + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + responseObserverRef.set(responseObserver); + streamActiveLatch.countDown(); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) {} + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicReference> wrappedCallRef = + new AtomicReference<>(); + ServerInterceptor capturingInterceptor = + new ServerInterceptor() { + @SuppressWarnings("unchecked") + @Override + public ServerCall.Listener interceptCall( + ServerCall call, Metadata headers, ServerCallHandler next) { + wrappedCallRef.set((ServerCall) call); + return next.startCall(call, headers); + } + }; + + final AtomicReference rawCallRef = new AtomicReference<>(); + ServerInterceptor capturingRawCallInterceptor = + new ServerInterceptor() { + @SuppressWarnings("unchecked") + @Override + public ServerCall.Listener interceptCall( + ServerCall call, Metadata headers, ServerCallHandler next) { + ConcurrencyDetectingServerCall wrapped = + new ConcurrencyDetectingServerCall((ServerCall) call); + rawCallRef.set(wrapped); + return next.startCall((ServerCall) (ServerCall) wrapped, headers); + } + }; + + startDataPlane(capturingRawCallInterceptor, interceptor, capturingInterceptor); + + Metadata initialHeaders = new Metadata(); + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + initialHeaders); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + boolean active = streamActiveLatch.await(5, TimeUnit.SECONDS); + assertThat(active).isTrue(); + + StreamObserver responseObserver = responseObserverRef.get(); + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders( + HeadersResponse.newBuilder() + .setResponse(CommonResponse.newBuilder().build()) + .build()) + .build()); + + ServerCall wrappedCall = wrappedCallRef.get(); + assertThat(wrappedCall).isNotNull(); + wrappedCall.sendHeaders(new Metadata()); + + ConcurrencyDetectingServerCall rawCall = rawCallRef.get(); + assertThat(rawCall).isNotNull(); + + responseObserver.onError(new RuntimeException("Stream dropped")); + + final CountDownLatch messageSentLatch = new CountDownLatch(2); + Thread appThread = + new Thread( + () -> { + try { + wrappedCall.sendMessage( + new ByteArrayInputStream("app-msg".getBytes(StandardCharsets.UTF_8))); + messageSentLatch.countDown(); + } catch (Exception e) { + // ignore + } + }); + + wrappedCall.sendMessage( + new ByteArrayInputStream("buffered-msg".getBytes(StandardCharsets.UTF_8))); + messageSentLatch.countDown(); + + appThread.start(); + appThread.join(); + + boolean sent = messageSentLatch.await(5, TimeUnit.SECONDS); + assertThat(sent).isTrue(); + assertThat(rawCall.concurrentCallDetected.get()).isFalse(); + } + + @Test + public void serverInterceptor_concurrentSendHeadersAndFailOpen_flushesHeadersCorrectly() + throws Exception { + ExternalProcessor proto = createBaseProto(extProcServerName).setFailureModeAllow(true).build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final AtomicReference> responseObserverRef = + new AtomicReference<>(); + final CountDownLatch streamActiveLatch = new CountDownLatch(1); + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + responseObserverRef.set(responseObserver); + streamActiveLatch.countDown(); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) {} + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicReference> wrappedCallRef = + new AtomicReference<>(); + ServerInterceptor capturingInterceptor = + new ServerInterceptor() { + @SuppressWarnings("unchecked") + @Override + public ServerCall.Listener interceptCall( + ServerCall call, Metadata headers, ServerCallHandler next) { + wrappedCallRef.set((ServerCall) call); + return next.startCall(call, headers); + } + }; + + final AtomicReference rawCallRef = new AtomicReference<>(); + ServerInterceptor capturingRawCallInterceptor = + new ServerInterceptor() { + @SuppressWarnings("unchecked") + @Override + public ServerCall.Listener interceptCall( + ServerCall call, Metadata headers, ServerCallHandler next) { + ConcurrencyDetectingServerCall wrapped = + new ConcurrencyDetectingServerCall((ServerCall) call); + rawCallRef.set(wrapped); + return next.startCall((ServerCall) (ServerCall) wrapped, headers); + } + }; + + startDataPlane(capturingRawCallInterceptor, interceptor, capturingInterceptor); + + Metadata initialHeaders = new Metadata(); + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch headersReceivedLatch = new CountDownLatch(1); + final AtomicReference receivedResponseHeadersRef = new AtomicReference<>(); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onHeaders(Metadata headers) { + receivedResponseHeadersRef.set(headers); + headersReceivedLatch.countDown(); + } + }, + initialHeaders); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + boolean active = streamActiveLatch.await(5, TimeUnit.SECONDS); + assertThat(active).isTrue(); + + StreamObserver responseObserver = responseObserverRef.get(); + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders( + HeadersResponse.newBuilder() + .setResponse(CommonResponse.newBuilder().build()) + .build()) + .build()); + + ServerCall wrappedCall = wrappedCallRef.get(); + assertThat(wrappedCall).isNotNull(); + + ConcurrencyDetectingServerCall rawCall = rawCallRef.get(); + assertThat(rawCall).isNotNull(); + + Thread appThread = + new Thread( + () -> { + Metadata headers = new Metadata(); + headers.put( + Metadata.Key.of("x-resp-header", Metadata.ASCII_STRING_MARSHALLER), "val"); + wrappedCall.sendHeaders(headers); + wrappedCall.close(Status.OK, new Metadata()); + }); + + appThread.start(); + responseObserver.onError(new RuntimeException("Stream failure")); + + appThread.join(); + + boolean headersSent = headersReceivedLatch.await(5, TimeUnit.SECONDS); + assertThat(headersSent).isTrue(); + assertThat( + receivedResponseHeadersRef + .get() + .get(Metadata.Key.of("x-resp-header", Metadata.ASCII_STRING_MARSHALLER))) + .isEqualTo("val"); + assertThat(rawCall.concurrentCallDetected.get()).isFalse(); + } + + // Category 26: Request-Scoped Context Propagation + @Test + public void serverInterceptor_contextPropagatedToStartCall() throws Exception { + ExternalProcessor proto = createBaseProto(extProcServerName).build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = new CountDownLatch(1); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders( + HeadersResponse.newBuilder() + .setResponse(CommonResponse.newBuilder().build()) + .build()) + .build()); + responseObserver.onCompleted(); + extProcLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + ServerInterceptor contextAttachingInterceptor = + new ServerInterceptor() { + @Override + public ServerCall.Listener interceptCall( + ServerCall call, Metadata headers, ServerCallHandler next) { + Context contextWithKey = Context.current().withValue(TRACE_KEY, "test-trace-123"); + return Contexts.interceptCall(contextWithKey, call, headers, next); + } + }; + + final AtomicBoolean startCallVerified = new AtomicBoolean(false); + + ServerInterceptor capturingInterceptor = + new ServerInterceptor() { + @Override + public ServerCall.Listener interceptCall( + ServerCall call, Metadata headers, ServerCallHandler next) { + if ("test-trace-123".equals(TRACE_KEY.get())) { + startCallVerified.set(true); + } + return next.startCall(call, headers); + } + }; + + startDataPlane(capturingInterceptor, interceptor, contextAttachingInterceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + assertThat(startCallVerified.get()).isTrue(); + } + + @Test + public void serverInterceptor_contextPropagatedToListenerCallbacks() throws Exception { + ExternalProcessor proto = createBaseProto(extProcServerName) + .setFailureModeAllow(true) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = new CountDownLatch(1); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders( + HeadersResponse.newBuilder() + .setResponse(CommonResponse.newBuilder().build()) + .build()) + .build()); + responseObserver.onCompleted(); + extProcLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + ServerInterceptor contextAttachingInterceptor = + new ServerInterceptor() { + @Override + public ServerCall.Listener interceptCall( + ServerCall call, Metadata headers, ServerCallHandler next) { + Context contextWithKey = Context.current().withValue(TRACE_KEY, "test-trace-123"); + return Contexts.interceptCall(contextWithKey, call, headers, next); + } + }; + + final AtomicBoolean onMessageVerified = new AtomicBoolean(false); + final AtomicBoolean onHalfCloseVerified = new AtomicBoolean(false); + final AtomicBoolean onReadyVerified = new AtomicBoolean(false); + + ServerInterceptor capturingInterceptor = + new ServerInterceptor() { + @Override + public ServerCall.Listener interceptCall( + ServerCall call, Metadata headers, ServerCallHandler next) { + ServerCall.Listener nextListener = next.startCall(call, headers); + return new ServerCall.Listener() { + @Override + public void onMessage(ReqT message) { + if ("test-trace-123".equals(TRACE_KEY.get())) { + onMessageVerified.set(true); + } + nextListener.onMessage(message); + } + + @Override + public void onHalfClose() { + if ("test-trace-123".equals(TRACE_KEY.get())) { + onHalfCloseVerified.set(true); + } + nextListener.onHalfClose(); + } + + @Override + public void onReady() { + if ("test-trace-123".equals(TRACE_KEY.get())) { + onReadyVerified.set(true); + } + nextListener.onReady(); + } + }; + } + }; + + startDataPlane(capturingInterceptor, interceptor, contextAttachingInterceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + assertThat(onMessageVerified.get()).isTrue(); + assertThat(onHalfCloseVerified.get()).isTrue(); + assertThat(onReadyVerified.get()).isTrue(); + } + + @Test + public void serverInterceptor_contextPropagatedToExtProcStub() throws Exception { + ExternalProcessor proto = createBaseProto(extProcServerName).build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = new CountDownLatch(1); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + @SuppressWarnings("unchecked") + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders( + HeadersResponse.newBuilder() + .setResponse(CommonResponse.newBuilder().build()) + .build()) + .build()); + responseObserver.onCompleted(); + extProcLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + final AtomicBoolean extProcStubContextVerified = new AtomicBoolean(false); + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName) + .directExecutor() + .intercept( + new io.grpc.ClientInterceptor() { + @Override + public io.grpc.ClientCall interceptCall( + MethodDescriptor method, + io.grpc.CallOptions callOptions, + io.grpc.Channel next) { + if (method + .getFullMethodName() + .equals( + ExternalProcessorGrpc.getProcessMethod() + .getFullMethodName())) { + if ("test-trace-123".equals(TRACE_KEY.get())) { + extProcStubContextVerified.set(true); + } + } + return next.newCall(method, callOptions); + } + }) + .build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + ServerInterceptor contextAttachingInterceptor = + new ServerInterceptor() { + @Override + public ServerCall.Listener interceptCall( + ServerCall call, Metadata headers, ServerCallHandler next) { + Context contextWithKey = Context.current().withValue(TRACE_KEY, "test-trace-123"); + return Contexts.interceptCall(contextWithKey, call, headers, next); + } + }; + + ServerInterceptor capturingInterceptor = + new ServerInterceptor() { + @Override + public ServerCall.Listener interceptCall( + ServerCall call, Metadata headers, ServerCallHandler next) { + return next.startCall(call, headers); + } + }; + + startDataPlane(capturingInterceptor, interceptor, contextAttachingInterceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + assertThat(extProcStubContextVerified.get()).isTrue(); + } + + @Test + public void serialization_specCompliance() throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SEND) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = new CountDownLatch(1); + final AtomicReference capturedRequest = new AtomicReference<>(); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasResponseHeaders()) { + capturedRequest.set(request); + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders(HeadersResponse.newBuilder().build()) + .build()); + extProcLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + ServerInterceptor headersInterceptor = + new ServerInterceptor() { + @Override + public ServerCall.Listener interceptCall( + ServerCall call, Metadata headers, ServerCallHandler next) { + return next.startCall( + new ForwardingServerCall.SimpleForwardingServerCall(call) { + @Override + public void sendHeaders(Metadata responseHeaders) { + responseHeaders.put( + Metadata.Key.of("custom-ascii", Metadata.ASCII_STRING_MARSHALLER), + "hello-world"); + responseHeaders.put( + Metadata.Key.of("custom-bin", Metadata.BINARY_BYTE_MARSHALLER), + new byte[] {0x00, 0x01, 0x02}); + super.sendHeaders(responseHeaders); + } + }, + headers); + } + }; + + startDataPlane(headersInterceptor, interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + ProcessingRequest req = capturedRequest.get(); + assertThat(req).isNotNull(); + + io.envoyproxy.envoy.config.core.v3.HeaderMap headerMap = req.getResponseHeaders().getHeaders(); + io.envoyproxy.envoy.config.core.v3.HeaderValue customAsciiProto = null; + io.envoyproxy.envoy.config.core.v3.HeaderValue customBinProto = null; + for (io.envoyproxy.envoy.config.core.v3.HeaderValue hv : headerMap.getHeadersList()) { + if (hv.getKey().equals("custom-ascii")) { + customAsciiProto = hv; + } else if (hv.getKey().equals("custom-bin")) { + customBinProto = hv; + } + } + + assertThat(customAsciiProto).isNotNull(); + assertThat(customAsciiProto.getValue()).isEmpty(); + assertThat(customAsciiProto.getRawValue().toStringUtf8()).isEqualTo("hello-world"); + + assertThat(customBinProto).isNotNull(); + assertThat(customBinProto.getValue()).isEmpty(); + String expectedBase64 = + com.google.common.io.BaseEncoding.base64().encode(new byte[] {0x00, 0x01, 0x02}); + assertThat(customBinProto.getRawValue().toStringUtf8()).isEqualTo(expectedBase64); + } + + @Test + public void deserialization_preferRawValue() throws Exception { + ExternalProcessor proto = createBaseProto(extProcServerName).build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = new CountDownLatch(2); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + extProcLatch.countDown(); + } else if (request.hasResponseHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders( + HeadersResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setHeaderMutation( + HeaderMutation.newBuilder() + .addSetHeaders( + io.envoyproxy.envoy.config.core.v3 + .HeaderValueOption.newBuilder() + .setHeader( + io.envoyproxy.envoy.config.core.v3 + .HeaderValue.newBuilder() + .setKey("custom-ascii") + .setValue("legacy-val") + .setRawValue( + ByteString.copyFromUtf8( + "raw-val")) + .build()) + .build()) + .build()) + .build()) + .build()) + .build()); + extProcLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final AtomicReference receivedResponseHeaders = new AtomicReference<>(); + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onHeaders(Metadata headers) { + receivedResponseHeaders.set(headers); + } + + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(receivedResponseHeaders.get()).isNotNull(); + assertThat( + receivedResponseHeaders + .get() + .get(Metadata.Key.of("custom-ascii", Metadata.ASCII_STRING_MARSHALLER))) + .isEqualTo("raw-val"); + } + + @Test + public void deserialization_binaryHeader_validBase64() throws Exception { + ExternalProcessor proto = createBaseProto(extProcServerName).build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = new CountDownLatch(2); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + extProcLatch.countDown(); + } else if (request.hasResponseHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders( + HeadersResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setHeaderMutation( + HeaderMutation.newBuilder() + .addSetHeaders( + io.envoyproxy.envoy.config.core.v3 + .HeaderValueOption.newBuilder() + .setHeader( + io.envoyproxy.envoy.config.core.v3 + .HeaderValue.newBuilder() + .setKey("custom-bin") + .setRawValue( + ByteString.copyFromUtf8( + "YmFy")) + .build()) + .build()) + .build()) + .build()) + .build()) + .build()); + extProcLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final AtomicReference receivedResponseHeaders = new AtomicReference<>(); + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onHeaders(Metadata headers) { + receivedResponseHeaders.set(headers); + } + + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(receivedResponseHeaders.get()).isNotNull(); + byte[] binValue = + receivedResponseHeaders + .get() + .get(Metadata.Key.of("custom-bin", Metadata.BINARY_BYTE_MARSHALLER)); + assertThat(binValue).isEqualTo(new byte[] {'b', 'a', 'r'}); + } + + @Test + public void deserialization_binaryHeader_invalidBase64_noError_fails() throws Exception { + ExternalProcessor proto = createBaseProto(extProcServerName).build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = new CountDownLatch(2); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + extProcLatch.countDown(); + } else if (request.hasResponseHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders( + HeadersResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setHeaderMutation( + HeaderMutation.newBuilder() + .addSetHeaders( + io.envoyproxy.envoy.config.core.v3 + .HeaderValueOption.newBuilder() + .setHeader( + io.envoyproxy.envoy.config.core.v3 + .HeaderValue.newBuilder() + .setKey("custom-bin") + .setRawValue( + ByteString.copyFromUtf8( + "invalid_base64!")) + .build()) + .build()) + .build()) + .build()) + .build()) + .build()); + extProcLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final AtomicReference callStatus = new AtomicReference<>(); + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callStatus.set(status); + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callStatus.get()).isNotNull(); + assertThat(callStatus.get().getCode()).isEqualTo(Status.Code.INTERNAL); + assertThat(callStatus.get().getDescription()).contains("External processor stream failed"); + } + + @Test + public void deserialization_binaryHeader_invalidBase64_failsCall() throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setMutationRules( + io.envoyproxy.envoy.config.common.mutation_rules.v3.HeaderMutationRules.newBuilder() + .setDisallowIsError(com.google.protobuf.BoolValue.of(true)) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = new CountDownLatch(2); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + extProcLatch.countDown(); + } else if (request.hasResponseHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders( + HeadersResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setHeaderMutation( + HeaderMutation.newBuilder() + .addSetHeaders( + io.envoyproxy.envoy.config.core.v3 + .HeaderValueOption.newBuilder() + .setHeader( + io.envoyproxy.envoy.config.core.v3 + .HeaderValue.newBuilder() + .setKey("custom-bin") + .setRawValue( + ByteString.copyFromUtf8( + "invalid_base64!")) + .build()) + .build()) + .build()) + .build()) + .build()) + .build()); + extProcLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final AtomicReference callStatus = new AtomicReference<>(); + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callStatus.set(status); + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callStatus.get()).isNotNull(); + assertThat(callStatus.get().getCode()).isEqualTo(Status.Code.INTERNAL); + assertThat(callStatus.get().getDescription()).contains("External processor stream failed"); + } + + @Test + public void deserialization_asciiHeader_invalidChars_noError_fails() throws Exception { + ExternalProcessor proto = createBaseProto(extProcServerName).build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = new CountDownLatch(2); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + extProcLatch.countDown(); + } else if (request.hasResponseHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders( + HeadersResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setHeaderMutation( + HeaderMutation.newBuilder() + .addSetHeaders( + io.envoyproxy.envoy.config.core.v3 + .HeaderValueOption.newBuilder() + .setHeader( + io.envoyproxy.envoy.config.core.v3 + .HeaderValue.newBuilder() + .setKey("custom-ascii") + .setRawValue( + ByteString.copyFromUtf8( + "value_with_newline\n")) + .build()) + .build()) + .build()) + .build()) + .build()) + .build()); + extProcLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final AtomicReference callStatus = new AtomicReference<>(); + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callStatus.set(status); + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callStatus.get()).isNotNull(); + assertThat(callStatus.get().getCode()).isEqualTo(Status.Code.INTERNAL); + assertThat(callStatus.get().getDescription()).contains("External processor stream failed"); + } + + @Test + public void deserialization_asciiHeader_invalidCharacters_failsCall() throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setMutationRules( + io.envoyproxy.envoy.config.common.mutation_rules.v3.HeaderMutationRules.newBuilder() + .setDisallowIsError(com.google.protobuf.BoolValue.of(true)) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = new CountDownLatch(2); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + extProcLatch.countDown(); + } else if (request.hasResponseHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders( + HeadersResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setHeaderMutation( + HeaderMutation.newBuilder() + .addSetHeaders( + io.envoyproxy.envoy.config.core.v3 + .HeaderValueOption.newBuilder() + .setHeader( + io.envoyproxy.envoy.config.core.v3 + .HeaderValue.newBuilder() + .setKey("custom-ascii") + .setRawValue( + ByteString.copyFromUtf8( + "value_with_newline\n")) + .build()) + .build()) + .build()) + .build()) + .build()) + .build()); + extProcLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final AtomicReference callStatus = new AtomicReference<>(); + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callStatus.set(status); + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callStatus.get()).isNotNull(); + assertThat(callStatus.get().getCode()).isEqualTo(Status.Code.INTERNAL); + assertThat(callStatus.get().getDescription()).contains("External processor stream failed"); + } + + @Test + public void deserialization_headerValue_tooLong_noError_fails() throws Exception { + ExternalProcessor proto = createBaseProto(extProcServerName).build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = new CountDownLatch(2); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + extProcLatch.countDown(); + } else if (request.hasResponseHeaders()) { + String longVal = com.google.common.base.Strings.repeat("a", 16385); + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders( + HeadersResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setHeaderMutation( + HeaderMutation.newBuilder() + .addSetHeaders( + io.envoyproxy.envoy.config.core.v3 + .HeaderValueOption.newBuilder() + .setHeader( + io.envoyproxy.envoy.config.core.v3 + .HeaderValue.newBuilder() + .setKey("custom-ascii") + .setRawValue( + ByteString.copyFromUtf8( + longVal)) + .build()) + .build()) + .build()) + .build()) + .build()) + .build()); + extProcLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final AtomicReference callStatus = new AtomicReference<>(); + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callStatus.set(status); + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callStatus.get()).isNotNull(); + assertThat(callStatus.get().getCode()).isEqualTo(Status.Code.INTERNAL); + assertThat(callStatus.get().getDescription()).contains("External processor stream failed"); + } + + @Test + public void deserialization_headerValue_tooLong_failsCall() throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setMutationRules( + io.envoyproxy.envoy.config.common.mutation_rules.v3.HeaderMutationRules.newBuilder() + .setDisallowIsError(com.google.protobuf.BoolValue.of(true)) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = new CountDownLatch(2); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + extProcLatch.countDown(); + } else if (request.hasResponseHeaders()) { + String longVal = com.google.common.base.Strings.repeat("a", 16385); + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseHeaders( + HeadersResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setHeaderMutation( + HeaderMutation.newBuilder() + .addSetHeaders( + io.envoyproxy.envoy.config.core.v3 + .HeaderValueOption.newBuilder() + .setHeader( + io.envoyproxy.envoy.config.core.v3 + .HeaderValue.newBuilder() + .setKey("custom-ascii") + .setRawValue( + ByteString.copyFromUtf8( + longVal)) + .build()) + .build()) + .build()) + .build()) + .build()) + .build()); + extProcLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final AtomicReference callStatus = new AtomicReference<>(); + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callStatus.set(status); + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callStatus.get()).isNotNull(); + assertThat(callStatus.get().getCode()).isEqualTo(Status.Code.INTERNAL); + assertThat(callStatus.get().getDescription()).contains("External processor stream failed"); + } + + @Test + public void givenRequestBodyModeGrpc_whenClientSendsEmptyMessage_thenEmptyMessageIsDelivered() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SKIP) + .setRequestBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseBodyMode(ProcessingMode.BodySendMode.NONE) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = new CountDownLatch(1); + final AtomicReference capturedRequest = new AtomicReference<>(); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestBody()) { + capturedRequest.set(request); + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setRequestBody( + BodyResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setEndOfStream(true) + .setEndOfStreamWithoutMessage( + request + .getRequestBody() + .getEndOfStreamWithoutMessage()) + .build()) + .build()) + .build()) + .build()) + .build()); + extProcLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicReference receivedRequest = new AtomicReference<>(); + final CountDownLatch dataPlaneLatch = new CountDownLatch(1); + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + try { + receivedRequest.set(request); + responseObserver.onNext( + new ByteArrayInputStream("Hello".getBytes(StandardCharsets.UTF_8))); + responseObserver.onCompleted(); + dataPlaneLatch.countDown(); + } catch (Throwable t) { + responseObserver.onError(t); + } + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream(new byte[0])); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(dataPlaneLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + ProcessingRequest req = capturedRequest.get(); + assertThat(req).isNotNull(); + assertThat(req.getRequestBody().getBody().isEmpty()).isTrue(); + assertThat(req.getRequestBody().getEndOfStreamWithoutMessage()).isFalse(); + + InputStream serverReceivedStream = receivedRequest.get(); + assertThat(serverReceivedStream).isNotNull(); + assertThat(serverReceivedStream.available()).isEqualTo(0); + } + + @Test + public void givenResponseBodyModeGrpc_whenServerSendsEmptyMessage_thenEmptyMessageIsDelivered() + throws Exception { + ExternalProcessor proto = + createBaseProto(extProcServerName) + .setProcessingMode( + ProcessingMode.newBuilder() + .setRequestHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SEND) + .setRequestBodyMode(ProcessingMode.BodySendMode.NONE) + .setResponseBodyMode(ProcessingMode.BodySendMode.GRPC) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch extProcLatch = new CountDownLatch(1); + final AtomicReference capturedRequest = new AtomicReference<>(); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasResponseBody()) { + capturedRequest.set(request); + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseBody( + BodyResponse.newBuilder() + .setResponse( + CommonResponse.newBuilder() + .setBodyMutation( + BodyMutation.newBuilder() + .setStreamedResponse( + StreamedBodyResponse.newBuilder() + .setBody(ByteString.EMPTY) + .setEndOfStream( + request + .getResponseBody() + .getEndOfStream()) + .build()) + .build()) + .build()) + .build()) + .build()); + extProcLatch.countDown(); + } else if (request.hasResponseTrailers()) { + responseObserver.onNext( + ProcessingResponse.newBuilder() + .setResponseTrailers(TrailersResponse.newBuilder().build()) + .build()); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + }; + + grpcCleanup.register( + InProcessServerBuilder.forName(extProcServerName) + .addService(extProcImpl) + .directExecutor() + .build() + .start()); + + CachedChannelManager channelManager = + new CachedChannelManager( + config -> + grpcCleanup.register( + InProcessChannelBuilder.forName(extProcServerName).directExecutor().build())); + + ExternalProcessorServerInterceptor interceptor = + new ExternalProcessorServerInterceptor(filterConfig, channelManager, FAKE_CONTEXT); + + dataPlaneHandler = + new DataPlaneServiceHandler() { + @Override + public void sayHello(InputStream request, StreamObserver responseObserver) { + responseObserver.onNext( + new ByteArrayInputStream(new byte[0])); // Empty response message! + responseObserver.onCompleted(); + } + }; + + startDataPlane(interceptor); + + io.grpc.ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final AtomicReference receivedResponse = new AtomicReference<>(); + final CountDownLatch clientLatch = new CountDownLatch(1); + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start( + new io.grpc.ClientCall.Listener() { + @Override + public void onMessage(InputStream message) { + receivedResponse.set(message); + clientLatch.countDown(); + } + + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, + new Metadata()); + + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(clientLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + ProcessingRequest req = capturedRequest.get(); + assertThat(req).isNotNull(); + assertThat(req.getResponseBody().getBody().isEmpty()).isTrue(); + + InputStream clientReceivedStream = receivedResponse.get(); + assertThat(clientReceivedStream).isNotNull(); + assertThat(clientReceivedStream.available()).isEqualTo(0); + } + + @Test + public void testFlowControlStateInitialization() throws Exception { + ExternalProcessor proto = createBaseProto(extProcServerName) + .setProcessingMode(ProcessingMode.newBuilder() + .setRequestBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SEND) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SEND) + .build()) + .build(); + + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final List receivedRequests = new java.util.concurrent.CopyOnWriteArrayList<>(); + final CountDownLatch extProcLatch = new CountDownLatch(2); + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + final StreamObserver responseObserver) { + System.out.println("extProcImpl.process called"); + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + System.out.println("extProcImpl.onNext:\n" + request); + receivedRequests.add(request); + extProcLatch.countDown(); + if (request.hasRequestHeaders()) { + System.out.println("extProcImpl: responding to request headers"); + responseObserver.onNext(ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + } else if (request.hasRequestBody()) { + boolean isEOS = request.getRequestBody().getEndOfStream(); + boolean isEOSWithoutMsg = request.getRequestBody().getEndOfStreamWithoutMessage(); + if (isEOSWithoutMsg || (isEOS && request.getRequestBody().getBody().isEmpty())) { + System.out.println("extProcImpl: sending EOS without message"); + responseObserver.onNext(ProcessingResponse.newBuilder() + .setRequestBody(BodyResponse.newBuilder() + .setResponse(CommonResponse.newBuilder() + .setBodyMutation(BodyMutation.newBuilder() + .setStreamedResponse(StreamedBodyResponse.newBuilder() + .setEndOfStreamWithoutMessage(true) + .build()) + .build()) + .build()) + .build()) + .build()); + } else { + System.out.println("extProcImpl: sending body response, EOS=" + isEOS); + responseObserver.onNext(ProcessingResponse.newBuilder() + .setRequestBody(BodyResponse.newBuilder() + .setResponse(CommonResponse.newBuilder() + .setBodyMutation(BodyMutation.newBuilder() + .setStreamedResponse(StreamedBodyResponse.newBuilder() + .setBody(request.getRequestBody().getBody()) + .setEndOfStream(isEOS) + .build()) + .build()) + .build()) + .build()) + .build()); + } + } else if (request.hasResponseHeaders()) { + System.out.println("extProcImpl: responding to response headers"); + responseObserver.onNext(ProcessingResponse.newBuilder() + .setResponseHeaders(HeadersResponse.newBuilder().build()) + .build()); + } else if (request.hasResponseBody()) { + System.out.println("extProcImpl: responding to response body"); + responseObserver.onNext(ProcessingResponse.newBuilder() + .setResponseBody(BodyResponse.newBuilder() + .setResponse(CommonResponse.newBuilder() + .setBodyMutation(BodyMutation.newBuilder() + .setStreamedResponse(StreamedBodyResponse.newBuilder() + .setBody(request.getResponseBody().getBody()) + .build()) + .build()) + .build()) + .build()) + .build()); + } else if (request.hasResponseTrailers()) { + System.out.println("extProcImpl: responding to response trailers"); + responseObserver.onNext(ProcessingResponse.newBuilder() + .setResponseTrailers(TrailersResponse.newBuilder().build()) + .build()); + } + } + + @Override + public void onError(Throwable t) { + System.out.println("extProcImpl.onError: " + t); + } + + @Override + public void onCompleted() { + System.out.println("extProcImpl.onCompleted"); + responseObserver.onCompleted(); + } + }; + } + }; + + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + grpcCleanup.register(InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build().start()); + + CachedChannelManager channelManager = new CachedChannelManager(config -> { + return grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName).directExecutor().build()); + }); + + ExternalProcessorServerInterceptor interceptor = new ExternalProcessorServerInterceptor( + filterConfig, channelManager, FAKE_CONTEXT); + + startDataPlane(interceptor); + + ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_RAW, io.grpc.CallOptions.DEFAULT); + + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + clientCall.start(new ClientCall.Listener() { + @Override + public void onHeaders(Metadata headers) { + System.out.println("clientCall.onHeaders"); + } + @Override + public void onMessage(InputStream message) { + System.out.println("clientCall.onMessage"); + } + @Override + public void onClose(Status status, Metadata trailers) { + System.out.println("clientCall.onClose: " + status); + callCompletedLatch.countDown(); + } + }, new Metadata()); + clientCall.request(1); + clientCall.sendMessage(new ByteArrayInputStream("hello".getBytes(StandardCharsets.UTF_8))); + clientCall.halfClose(); + + assertThat(extProcLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + channelManager.close(); + + assertThat(receivedRequests).hasSize(6); + ProcessingRequest firstRequest = receivedRequests.get(0); + ProcessingRequest secondRequest = receivedRequests.get(1); + ProcessingRequest thirdRequest = receivedRequests.get(2); + ProcessingRequest fourthRequest = receivedRequests.get(3); + ProcessingRequest fifthRequest = receivedRequests.get(4); + ProcessingRequest sixthRequest = receivedRequests.get(5); + + assertThat(firstRequest.hasRequestHeaders()).isTrue(); + assertThat(firstRequest.hasFlowControlInit()).isTrue(); + assertThat(firstRequest.getFlowControlInit().getInitialWindowDownstreamToSidestream()) + .isEqualTo(65536); + assertThat(firstRequest.getFlowControlInit().getInitialWindowSidestreamToUpstream()) + .isEqualTo(65536); + + assertThat(secondRequest.hasRequestBody()).isTrue(); + assertThat(secondRequest.hasFlowControlInit()).isFalse(); + + assertThat(thirdRequest.hasRequestBody()).isTrue(); + assertThat(thirdRequest.getRequestBody().getEndOfStreamWithoutMessage()).isTrue(); + assertThat(thirdRequest.hasClientWindowUpdate()).isTrue(); + assertThat(thirdRequest.getClientWindowUpdate().getWindowIncrementSidestreamToUpstream()) + .isEqualTo(5); + + assertThat(fourthRequest.hasResponseHeaders()).isTrue(); + assertThat(fifthRequest.hasResponseBody()).isTrue(); + assertThat(sixthRequest.hasResponseTrailers()).isTrue(); + } + + @Test + public void testDownstreamToSidestreamFlowControl_EnforcesWindow() throws Exception { + ExternalProcessor proto = createBaseProto(extProcServerName) + .setProcessingMode(ProcessingMode.newBuilder() + .setRequestBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseBodyMode(ProcessingMode.BodySendMode.NONE) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SKIP) + .build()) + .build(); + + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final List receivedRequests = new java.util.concurrent.CopyOnWriteArrayList<>(); + final CountDownLatch firstBodyLatch = new CountDownLatch(2); // Headers + First Body + final CountDownLatch secondBodyLatch = new CountDownLatch(1); + final AtomicReference> responseObserverRef = new AtomicReference<>(); + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + final StreamObserver responseObserver) { + responseObserverRef.set(responseObserver); + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + System.out.println("extProcImpl.onNext:\n" + request); + receivedRequests.add(request); + if (request.hasRequestHeaders()) { + responseObserver.onNext(ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + firstBodyLatch.countDown(); + } else if (request.hasRequestBody()) { + boolean isEOS = request.getRequestBody().getEndOfStream(); + boolean isEOSWithoutMsg = request.getRequestBody().getEndOfStreamWithoutMessage(); + if (isEOSWithoutMsg || (isEOS && request.getRequestBody().getBody().isEmpty())) { + responseObserver.onNext(ProcessingResponse.newBuilder() + .setRequestBody(BodyResponse.newBuilder() + .setResponse(CommonResponse.newBuilder() + .setBodyMutation(BodyMutation.newBuilder() + .setStreamedResponse(StreamedBodyResponse.newBuilder() + .setEndOfStreamWithoutMessage(true) + .build()) + .build()) + .build()) + .build()) + .build()); + } else { + responseObserver.onNext(ProcessingResponse.newBuilder() + .setRequestBody(BodyResponse.newBuilder() + .setResponse(CommonResponse.newBuilder() + .setBodyMutation(BodyMutation.newBuilder() + .setStreamedResponse(StreamedBodyResponse.newBuilder() + .setBody(request.getRequestBody().getBody()) + .setEndOfStream(isEOS) + .build()) + .build()) + .build()) + .build()) + .build()); + if (firstBodyLatch.getCount() > 0) { + firstBodyLatch.countDown(); + } else { + secondBodyLatch.countDown(); + } + } + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + grpcCleanup.register(InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build().start()); + + CachedChannelManager channelManager = new CachedChannelManager(config -> { + return grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName).directExecutor().build()); + }); + + ExternalProcessorServerInterceptor interceptor = new ExternalProcessorServerInterceptor( + filterConfig, channelManager, FAKE_CONTEXT); + + final List dataPlaneReceivedMessages = new java.util.concurrent.CopyOnWriteArrayList<>(); + final CountDownLatch serverAppCompletedLatch = new CountDownLatch(1); + + dataPlaneHandler = new DataPlaneServiceHandler() { + @Override + public StreamObserver sayHelloClientStreaming( + StreamObserver responseObserver) { + return new StreamObserver() { + @Override + public void onNext(InputStream value) { + try { + dataPlaneReceivedMessages.add(ByteString.readFrom(value)); + } catch (IOException e) { + throw new RuntimeException(e); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onNext(new ByteArrayInputStream("Response".getBytes(StandardCharsets.UTF_8))); + responseObserver.onCompleted(); + serverAppCompletedLatch.countDown(); + } + }; + } + }; + + startDataPlane(interceptor); + + ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_CLIENT_STREAMING, io.grpc.CallOptions.DEFAULT); + + final List clientResponseMessages = new java.util.concurrent.CopyOnWriteArrayList<>(); + final CountDownLatch callCompletedLatch = new CountDownLatch(1); + + clientCall.start(new ClientCall.Listener() { + @Override + public void onMessage(InputStream message) { + try { + clientResponseMessages.add(ByteString.readFrom(message)); + } catch (IOException e) { + throw new RuntimeException(e); + } + } + + @Override + public void onClose(Status status, Metadata trailers) { + callCompletedLatch.countDown(); + } + }, new Metadata()); + clientCall.request(1); + + // Generate large messages + ByteString largeMessage70k = ByteString.copyFrom(new byte[70000]); + ByteString largeMessage30k = ByteString.copyFrom(new byte[30000]); + + // Send first message (70000 bytes) - fits in 65536 window (sent immediately) + clientCall.sendMessage(new ByteArrayInputStream(largeMessage70k.toByteArray())); + assertThat(firstBodyLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + // Send second message (30000 bytes) - total 100000 > 65536, should be blocked in transport + clientCall.sendMessage(new ByteArrayInputStream(largeMessage30k.toByteArray())); + + // Verify it is NOT delivered immediately + assertThat(receivedRequests).hasSize(3); + assertThat(receivedRequests.get(0).hasRequestHeaders()).isTrue(); + assertThat(receivedRequests.get(1).hasRequestBody()).isTrue(); + assertThat(receivedRequests.get(2).hasClientWindowUpdate()).isTrue(); + + // Now send ServerWindowUpdate from ext_proc to interceptor to increment window by 40000 + responseObserverRef.get().onNext(ProcessingResponse.newBuilder() + .setServerWindowUpdate(ProcessingResponse.ServerWindowUpdate.newBuilder() + .setWindowIncrementDownstreamToSidestream(40000) + .build()) + .build()); + + // The second body should now be flushed and received by ext_proc + assertThat(secondBodyLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + // Client half-closes + clientCall.halfClose(); + + // Wait for the server App to complete and call to close successfully. + assertThat(serverAppCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(callCompletedLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + assertThat(dataPlaneReceivedMessages) + .containsExactly(largeMessage70k, largeMessage30k).inOrder(); + assertThat(clientResponseMessages) + .containsExactly(ByteString.copyFromUtf8("Response")); + + } + + @Test + public void testFailOpen_DrainsInboundQueuesInOrder() throws Exception { + ExternalProcessor proto = createBaseProto(extProcServerName) + .setFailureModeAllow(true) + .setProcessingMode(ProcessingMode.newBuilder() + .setRequestBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseBodyMode(ProcessingMode.BodySendMode.NONE) + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SKIP) + .build()) + .build(); + + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final AtomicReference> responseObserverRef = new AtomicReference<>(); + final CountDownLatch extProcActiveLatch = new CountDownLatch(1); + + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + final StreamObserver responseObserver) { + responseObserverRef.set(responseObserver); + ((ServerCallStreamObserver) responseObserver).request(100); + extProcActiveLatch.countDown(); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext(ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .build()); + } + } + @Override + public void onError(Throwable t) {} + @Override + public void onCompleted() {} + }; + } + }; + + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + grpcCleanup.register(InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build().start()); + + CachedChannelManager channelManager = new CachedChannelManager(config -> { + return grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName).directExecutor().build()); + }); + + ExternalProcessorServerInterceptor interceptor = new ExternalProcessorServerInterceptor( + filterConfig, channelManager, FAKE_CONTEXT); + + final AtomicReference + serverCallRef = new AtomicReference<>(); + + ServerInterceptor captureInterceptor = new ServerInterceptor() { + @Override + public ServerCall.Listener interceptCall( + ServerCall call, + Metadata headers, + ServerCallHandler next) { + serverCallRef.set((ExternalProcessorServerInterceptor.DataPlaneServerCall) call); + return next.startCall(call, headers); + } + }; + + final List dataPlaneReceivedMessages = new java.util.concurrent.CopyOnWriteArrayList<>(); + final CountDownLatch appCompletedLatch = new CountDownLatch(1); + + dataPlaneHandler = new DataPlaneServiceHandler() { + @Override + public StreamObserver sayHelloBidi( + StreamObserver responseObserver) { + ServerCallStreamObserver serverCallObserver = + (ServerCallStreamObserver) responseObserver; + serverCallObserver.disableAutoRequest(); + return new StreamObserver() { + @Override + public void onNext(InputStream value) { + try { + dataPlaneReceivedMessages.add(ByteString.readFrom(value)); + } catch (IOException e) { + throw new RuntimeException(e); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + appCompletedLatch.countDown(); + } + }; + } + }; + + // Start data plane with both interceptors + dataPlaneServerName = InProcessServerBuilder.generateName(); + ServerServiceDefinition dataPlaneService = + ServerServiceDefinition.builder("test.TestService") + .addMethod( + METHOD_SAY_HELLO_BIDI, + io.grpc.stub.ServerCalls.asyncBidiStreamingCall( + responseObserver -> dataPlaneHandler.sayHelloBidi(responseObserver))) + .build(); + + grpcCleanup.register( + InProcessServerBuilder.forName(dataPlaneServerName) + .addService( + ServerInterceptors.intercept( + dataPlaneService, java.util.Arrays.asList(captureInterceptor, interceptor))) + .directExecutor() + .build() + .start()); + + dataPlaneChannel = grpcCleanup.register( + InProcessChannelBuilder.forName(dataPlaneServerName).directExecutor().build()); + + ClientCall clientCall = + dataPlaneChannel.newCall(METHOD_SAY_HELLO_BIDI, io.grpc.CallOptions.DEFAULT); + + clientCall.start(new ClientCall.Listener() { + @Override + public void onMessage(InputStream message) {} + @Override + public void onClose(Status status, Metadata trailers) {} + }, new Metadata()); + + // Wait for ExtProc stream to be active + assertThat(extProcActiveLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + ExternalProcessorServerInterceptor.DataPlaneServerCall dataPlaneCall = serverCallRef.get(); + assertThat(dataPlaneCall).isNotNull(); + + ExternalProcessorServerInterceptor.DataPlaneServerListener listener = + dataPlaneCall.getListener(); + assertThat(listener).isNotNull(); + + // Manually populate queues + synchronized (dataPlaneCall.streamLock) { + dataPlaneCall.pendingMutatedRequestBodies.add( + StreamedBodyResponse.newBuilder() + .setBody(ByteString.copyFromUtf8("mutated-1")) + .build()); + dataPlaneCall.pendingRequestBodyMessages.add(ByteString.copyFromUtf8("raw-waiting-2")); + listener.savedMessages.add(new ByteArrayInputStream("raw-saved-3".getBytes(StandardCharsets.UTF_8))); + } + + // Verify nothing received yet + assertThat(dataPlaneReceivedMessages).isEmpty(); + + // Fail the ext_proc stream to trigger fail-open + responseObserverRef.get().onError(new RuntimeException("ExtProc stream failed")); + + assertThat(dataPlaneReceivedMessages).containsExactly( + ByteString.copyFromUtf8("mutated-1"), + ByteString.copyFromUtf8("raw-waiting-2"), + ByteString.copyFromUtf8("raw-saved-3") + ).inOrder(); + + clientCall.halfClose(); + channelManager.close(); + } +} diff --git a/xds/third_party/envoy/src/main/proto/envoy/service/ext_proc/v3/external_processor.proto b/xds/third_party/envoy/src/main/proto/envoy/service/ext_proc/v3/external_processor.proto index 1c033c08d26..a02779033be 100644 --- a/xds/third_party/envoy/src/main/proto/envoy/service/ext_proc/v3/external_processor.proto +++ b/xds/third_party/envoy/src/main/proto/envoy/service/ext_proc/v3/external_processor.proto @@ -23,37 +23,31 @@ option (udpa.annotations.file_status).package_version_status = ACTIVE; // [#protodoc-title: External processing service] -// A service that can access and modify HTTP requests and responses -// as part of a filter chain. +// A service that can access and modify HTTP requests and responses as part of a filter chain. // The overall external processing protocol works like this: // // 1. The data plane sends to the service information about the HTTP request. -// 2. The service sends back a ProcessingResponse message that directs -// the data plane to either stop processing, continue without it, or send -// it the next chunk of the message body. -// 3. If so requested, the data plane sends the server the message body in -// chunks, or the entire body at once. In either case, the server may send -// back a ProcessingResponse for each message it receives, or wait for -// a certain amount of body chunks received before streaming back the -// ProcessingResponse messages. -// 4. If so requested, the data plane sends the server the HTTP trailers, -// and the server sends back a ProcessingResponse. -// 5. At this point, request processing is done, and we pick up again -// at step 1 when the data plane receives a response from the upstream -// server. -// 6. At any point above, if the server closes the gRPC stream cleanly, -// then the data plane proceeds without consulting the server. -// 7. At any point above, if the server closes the gRPC stream with an error, -// then the data plane returns a 500 error to the client, unless the filter -// was configured to ignore errors. +// 2. The service sends back a ``ProcessingResponse`` message that directs the data plane to either +// stop processing, continue without it, or send it the next chunk of the message body. +// 3. If so requested, the data plane sends the server the message body in chunks, or the entire +// body at once. In either case, the server may send back a ``ProcessingResponse`` for each +// message it receives, or wait for a certain amount of body chunks to be received before +// streaming back the ``ProcessingResponse`` messages. +// 4. If so requested, the data plane sends the server the HTTP trailers, and the server sends back +// a ``ProcessingResponse``. +// 5. At this point, request processing is done, and we pick up again at step 1 when the data plane +// receives a response from the upstream server. +// 6. At any point above, if the server closes the gRPC stream cleanly, then the data plane +// proceeds without consulting the server. +// 7. At any point above, if the server closes the gRPC stream with an error, then the data plane +// returns a ``500`` error to the client, unless the filter was configured to ignore errors. // -// In other words, the process is a request/response conversation, but -// using a gRPC stream to make it easier for the server to -// maintain state. +// In other words, the process is a request/response conversation, but using a gRPC stream to make +// it easier for the server to maintain state. service ExternalProcessor { // This begins the bidirectional stream that the data plane will use to // give the server control over what the filter does. The actual - // protocol is described by the ProcessingRequest and ProcessingResponse + // protocol is described by the ``ProcessingRequest`` and ``ProcessingResponse`` // messages below. rpc Process(stream ProcessingRequest) returns (stream ProcessingResponse) { } @@ -61,30 +55,98 @@ service ExternalProcessor { // This message specifies the filter protocol configurations which will be sent to the ext_proc // server in a :ref:`ProcessingRequest `. -// If the server does not support these protocol configurations, it may choose to close the gRPC stream. -// If the server supports these protocol configurations, it should respond based on the API specifications. +// If the server does not support these protocol configurations, it may choose to close the gRPC +// stream. If the server supports these protocol configurations, it should respond based on the +// API specifications. message ProtocolConfiguration { - // Specify the filter configuration :ref:`request_body_mode - // ` + // Specifies the filter configuration + // :ref:`request_body_mode `. envoy.extensions.filters.http.ext_proc.v3.ProcessingMode.BodySendMode request_body_mode = 1 [(validate.rules).enum = {defined_only: true}]; - // Specify the filter configuration :ref:`response_body_mode - // ` + // Specifies the filter configuration + // :ref:`response_body_mode `. envoy.extensions.filters.http.ext_proc.v3.ProcessingMode.BodySendMode response_body_mode = 2 [(validate.rules).enum = {defined_only: true}]; - // Specify the filter configuration :ref:`send_body_without_waiting_for_header_response - // ` - // If the client is waiting for a header response from the server, setting ``true`` means the client will send body to the server - // as they arrive. Setting ``false`` means the client will buffer the arrived data and not send it to the server immediately. + // Specifies the filter configuration + // :ref:`send_body_without_waiting_for_header_response `. + // If the client is waiting for a header response from the server, setting to ``true`` means the + // client will send the body to the server as it arrives. Setting to ``false`` means the client + // will buffer the arrived data and not send it to the server immediately. bool send_body_without_waiting_for_header_response = 3; } // This represents the different types of messages that the data plane can send // to an external processing server. -// [#next-free-field: 12] +// [#next-free-field: 14] message ProcessingRequest { + // Initial flow control window sizes for ``FULL_DUPLEX_STREAMED`` and + // ``GRPC`` body send modes. + // + // A sender starts with this amount of flow control window. Whenever + // it sends body data, it must decrement its flow control window by + // the number of bytes that it has sent. When its flow control + // window is less than or equal to the amount of body data it wishes + // to send, it may not send until it receives a window update causing + // its flow control window to be large enough. + // + // However, note that in ``GRPC`` body send mode, whenever the flow + // control window is greater than zero, a sender may send a single + // message, even if the size of that message exceeds the available flow + // control window. At that point, the flow control window will be negative + // and the sender must not send the next message until it becomes positive. + // + // Note that the initial size for the to-sidestream windows are set by + // the sender, not the receiver. This is because each sidestream may be + // routed to a different ext_proc server instance, but there is no + // connection-level handshake to set a default for that server + // instance, so the only alternative here would be to have the + // ext_proc server instance set this on a per-stream basis, which + // would require an additional round-trip and therefore hurt latency. + // This unfortunately means that the ext_proc server instance has a + // bit less control: as soon as it receives these initial values, it can + // immediately send a window update that reduces the window, but it + // must be prepared to handle any data that the sender has already sent. + // The initial sizes for the to-sidestream windows are generally + // expected to be in the range of 32K to 64K. + // + // In ``FULL_DUPLEX_STREAMED`` body send mode, for backward compatibility + // with existing ext_proc servers that do not support flow control, if + // the ext_proc server's first response does not include a window + // update, or if the data plane does not receive the first response + // from the ext_proc server within the configured timeout, then the + // data plane will assume that the ext_proc server does not support + // flow control, and it will proceed without it. + // + // [#not-implemented-hide:] + message FlowControlInit { + // Downstream-to-sidestream initial window size. + int64 initial_window_downstream_to_sidestream = 1; + + // Sidestream-to-upstream initial window size. + int64 initial_window_sidestream_to_upstream = 2; + + // Upstream-to-sidestream initial window size. + int64 initial_window_upstream_to_sidestreama = 3; + + // Sidestream-to-downstream initial window size. + int64 initial_window_sidestream_to_downstream = 4; + } + + // Flow control window update. Values may be positive or negative. The + // sender must immediately add these values to its flow control window, + // which governs how much data can be sent. + // + // [#not-implemented-hide:] + message ClientWindowUpdate { + // Window update for sidestream-to-upstream. + int64 window_increment_sidestream_to_upstream = 1; + + // Window update for sidestream-to-downstream. + int64 window_increment_sidestream_to_downstream = 2; + } + reserved 1; reserved "async_mode"; @@ -93,35 +155,33 @@ message ProcessingRequest { // ones are set for a particular HTTP request/response depend on the // processing mode. oneof request { - option (validate.required) = true; - // Information about the HTTP request headers, as well as peer info and additional // properties. Unless ``observability_mode`` is ``true``, the server must send back a - // HeaderResponse message, an ImmediateResponse message, or close the stream. + // ``HeaderResponse`` message, an ``ImmediateResponse`` message, or close the stream. HttpHeaders request_headers = 2; // Information about the HTTP response headers, as well as peer info and additional // properties. Unless ``observability_mode`` is ``true``, the server must send back a - // HeaderResponse message or close the stream. + // ``HeaderResponse`` message or close the stream. HttpHeaders response_headers = 3; - // A chunk of the HTTP request body. Unless ``observability_mode`` is true, the server must send back - // a BodyResponse message, an ImmediateResponse message, or close the stream. + // A chunk of the HTTP request body. Unless ``observability_mode`` is ``true``, the server must + // send back a ``BodyResponse`` message, an ``ImmediateResponse`` message, or close the stream. HttpBody request_body = 4; - // A chunk of the HTTP response body. Unless ``observability_mode`` is ``true``, the server must send back - // a BodyResponse message or close the stream. + // A chunk of the HTTP response body. Unless ``observability_mode`` is ``true``, the server must + // send back a ``BodyResponse`` message or close the stream. HttpBody response_body = 5; // The HTTP trailers for the request path. Unless ``observability_mode`` is ``true``, the server - // must send back a TrailerResponse message or close the stream. + // must send back a ``TrailerResponse`` message or close the stream. // // This message is only sent if the trailers processing mode is set to ``SEND`` and // the original downstream request has trailers. HttpTrailers request_trailers = 6; // The HTTP trailers for the response path. Unless ``observability_mode`` is ``true``, the server - // must send back a TrailerResponse message or close the stream. + // must send back a ``TrailerResponse`` message or close the stream. // // This message is only sent if the trailers processing mode is set to ``SEND`` and // the original upstream response has trailers. @@ -137,39 +197,75 @@ message ProcessingRequest { // :ref:`attributes ` supported in the data plane. map attributes = 9; - // Specify whether the filter that sent this request is running in :ref:`observability_mode - // ` - // and defaults to false. + // Specifies whether the filter that sent this request is running in + // :ref:`observability_mode `. // - // * A value of ``false`` indicates that the server must respond - // to this message by either sending back a matching ProcessingResponse message, - // or by closing the stream. + // * A value of ``false`` indicates that the server must respond to this message by either + // sending back a matching ``ProcessingResponse`` message, or by closing the stream. // * A value of ``true`` indicates that the server should not respond to this message, as any - // responses will be ignored. However, it may still close the stream to indicate that no more messages - // are needed. + // responses will be ignored. However, it may still close the stream to indicate that no more + // messages are needed. // + // Defaults to ``false``. bool observability_mode = 10; // Specify the filter protocol configurations to be sent to the server. // ``protocol_config`` is only encoded in the first ``ProcessingRequest`` message from the client to the server. ProtocolConfiguration protocol_config = 11; + + // Flow control initialization for ``FULL_DUPLEX_STREAMED`` and + // ``GRPC`` body send modes. + // + // Must be set in the initial message on the stream. Not used in + // subsequent messages. + // + // [#not-implemented-hide:] + FlowControlInit flow_control_init = 12; + + // Flow control updates for ``FULL_DUPLEX_STREAMED`` and ``GRPC`` body + // send modes. + // + // This message may be included in a response message that also + // populates one of the fields in the ``request`` oneof above, or it + // may be sent in a response message that does not set the + // ``request`` oneof. + // + // In ``FULL_DUPLEX_STREAMED`` body send mode, for backward + // compatibility with data planes that do not yet support flow control, + // the data plane must not send a message containing only this field + // (i.e., not setting the ``request`` oneof) unless the ext_proc server + // has sent a window update, thus indicating that it supports flow control. + // + // [#not-implemented-hide:] + ClientWindowUpdate client_window_update = 13; } // This represents the different types of messages the server may send back to the data plane -// when the ``observability_mode`` field in the received ProcessingRequest is set to false. +// when the ``observability_mode`` field in the received ``ProcessingRequest`` is set to ``false``. // // * If the corresponding ``BodySendMode`` in the // :ref:`processing_mode ` -// is not set to ``FULL_DUPLEX_STREAMED``, then for every received ProcessingRequest, -// the server must send back exactly one ProcessingResponse message. +// is not set to ``FULL_DUPLEX_STREAMED``, then for every received ``ProcessingRequest``, +// the server must send back exactly one ``ProcessingResponse`` message. // * If it is set to ``FULL_DUPLEX_STREAMED``, the server must follow the API defined -// for this mode to send the ProcessingResponse messages. -// [#next-free-field: 13] +// for this mode to send the ``ProcessingResponse`` messages. +// [#next-free-field: 14] message ProcessingResponse { + // Flow control window update. Values may be positive or negative. The + // sender must immediately add these values to its flow control window, + // which governs how much data can be sent. + // + // [#not-implemented-hide:] + message ServerWindowUpdate { + // Window update for downstream-to-sidestream. + int64 window_increment_downstream_to_sidestream = 1; + + // Window update for upstream-to-sidestream. + int64 window_increment_upstream_to_sidestream = 2; + } + // The response type that is sent by the server. oneof response { - option (validate.required) = true; - // The server must send back this message in response to a message with the // ``request_headers`` field set. HeadersResponse request_headers = 1; @@ -204,17 +300,19 @@ message ProcessingResponse { ImmediateResponse immediate_response = 7; // The server sends back this message to initiate or continue local response streaming. - // The server must initiate local response streaming with the ``headers_response`` in response to a ProcessingRequest - // with the ``request_headers`` only. - // The server may follow up with multiple messages containing ``body_response``. The server must indicate - // end of stream by setting ``end_of_stream`` to ``true`` in the ``headers_response`` + // The server must initiate local response streaming with the ``headers_response`` in response + // to a ``ProcessingRequest`` with the ``request_headers`` only. + // The server may follow up with multiple messages containing ``body_response``. The server must + // indicate end of stream by setting ``end_of_stream`` to ``true`` in the ``headers_response`` // or ``body_response`` message or by sending a ``trailers_response`` message. - // The client may send a ``request_body`` or ``request_trailers`` to the server depending on configuration. + // The client may send a ``request_body`` or ``request_trailers`` to the server depending on + // configuration. // The streaming local response can only be sent when the ``request_header_mode`` in the filter // :ref:`processing_mode ` - // is set to ``SEND``. The ext_proc server should not send StreamedImmediateResponse if it did not observe request headers, - // as it will result in the race with the upstream server response and reset of the client request. - // Presently only the FULL_DUPLEX_STREAMED or NONE body modes are supported. + // is set to ``SEND``. The ext_proc server should not send ``StreamedImmediateResponse`` if it + // did not observe request headers, as it will result in a race with the upstream server + // response and reset of the client request. + // Presently only the ``FULL_DUPLEX_STREAMED`` or ``NONE`` body modes are supported. StreamedImmediateResponse streamed_immediate_response = 11; } @@ -223,19 +321,17 @@ message ProcessingResponse { // field name(s) of the struct. google.protobuf.Struct dynamic_metadata = 8; - // Override how parts of the HTTP request and response are processed - // for the duration of this particular request/response only. Servers - // may use this to intelligently control how requests are processed - // based on the headers and other metadata that they see. - // This field is only applicable when servers responding to the header requests. - // If it is set in the response to the body or trailer requests, it will be ignored by the data plane. + // Override how parts of the HTTP request and response are processed for the duration of this + // particular request/response only. Servers may use this to intelligently control how requests + // are processed based on the headers and other metadata that they see. + // + // This field is only applicable when servers are responding to the header requests. If it is set + // in the response to the body or trailer requests, it will be ignored by the data plane. // It is also ignored by the data plane when the ext_proc filter config - // :ref:`allow_mode_override - // ` - // is set to false, or - // :ref:`send_body_without_waiting_for_header_response - // ` - // is set to true. + // :ref:`allow_mode_override ` + // is set to ``false``, or + // :ref:`send_body_without_waiting_for_header_response ` + // is set to ``true``. envoy.extensions.filters.http.ext_proc.v3.ProcessingMode mode_override = 9; // [#not-implemented-hide:] @@ -251,70 +347,87 @@ message ProcessingResponse { // client had already sent before it saw the ext_proc stream termination. bool request_drain = 12; - // When ext_proc server receives a request message, in case it needs more - // time to process the message, it sends back a ProcessingResponse message - // with a new timeout value. When the data plane receives this response - // message, it ignores other fields in the response, just stop the original - // timer, which has the timeout value specified in - // :ref:`message_timeout - // ` - // and start a new timer with this ``override_message_timeout`` value and keep the - // data plane ext_proc filter state machine intact. - // Has to be >= 1ms and <= - // :ref:`max_message_timeout ` - // Such message can be sent at most once in a particular data plane ext_proc filter processing state. - // To enable this API, one has to set ``max_message_timeout`` to a number >= 1ms. + // When the ext_proc server receives a request message and needs more time to process it, it + // sends back a ``ProcessingResponse`` message with a new timeout value. When the data plane + // receives this response message, it ignores other fields in the response, stops the original + // timer (which has the timeout value specified in + // :ref:`message_timeout `), + // and starts a new timer with this ``override_message_timeout`` value while keeping the data + // plane ext_proc filter state machine intact. + // + // The value must be >= 1ms and <= + // :ref:`max_message_timeout `. + // Such a message can be sent at most once in a particular data plane ext_proc filter processing + // state. To enable this API, ``max_message_timeout`` must be set to a value >= 1ms. google.protobuf.Duration override_message_timeout = 10; + + // Flow control updates for ``FULL_DUPLEX_STREAMED`` and ``GRPC`` body + // send modes. + // + // This message may be included in a response message that also + // populates one of the fields in the ``response`` oneof above, or it + // may be sent in a response message that does not set the + // ``response`` oneof. + // + // In ``FULL_DUPLEX_STREAMED`` body send mode, for backward + // compatibility with data planes that do not yet support flow control, + // the ext_proc server must not set this field unless the data plane + // sent initial window sizes in its initial message on the stream. + // Conversely, if the data plane did send initial window sizes in its + // initial message on the stream, the ext_proc server must send a + // window update immediately to let the data plane know that it also + // supports flow control. If the ext_proc server is sending a message + // immediately anyway (e.g., for a header or body chunk), it can include + // this field in that same message; otherwise, the ext_proc server must + // send a message containing only this field. + // + // [#not-implemented-hide:] + ServerWindowUpdate server_window_update = 13; } // The following are messages that are sent to the server. -// This message is sent to the external server when the HTTP request and responses +// This message is sent to the external server when the HTTP request and response headers // are first received. message HttpHeaders { - // The HTTP request headers. All header keys will be - // lower-cased, because HTTP header keys are case-insensitive. - // The header value is encoded in the + // The HTTP request headers. All header keys will be lower-cased, because HTTP header keys are + // case-insensitive. The header value is encoded in the // :ref:`raw_value ` field. config.core.v3.HeaderMap headers = 1; // [#not-implemented-hide:] - // This field is deprecated and not implemented. Attributes will be sent in - // the top-level :ref:`attributes ` field. map attributes = 2 [deprecated = true, (envoy.annotations.deprecated_at_minor_version) = "3.0"]; - // If ``true``, then there is no message body associated with this - // request or response. + // If ``true``, then there is no message body associated with this request or response. bool end_of_stream = 3; } -// This message is sent to the external server when the HTTP request and -// response bodies are received. +// This message is sent to the external server when the HTTP request and response bodies are +// received. message HttpBody { - // The contents of the body in the HTTP request/response. Note that in - // streaming mode multiple ``HttpBody`` messages may be sent. + // The contents of the body in the HTTP request/response. Note that in streaming mode multiple + // ``HttpBody`` messages may be sent. // - // In ``GRPC`` body send mode, a separate ``HttpBody`` message will be - // sent for each message in the gRPC stream. + // In ``GRPC`` body send mode, a separate ``HttpBody`` message will be sent for each message in + // the gRPC stream. bytes body = 1; - // If ``true``, this will be the last ``HttpBody`` message that will be sent and no - // trailers will be sent for the current request/response. + // If ``true``, this will be the last ``HttpBody`` message that will be sent and no trailers + // will be sent for the current request/response. bool end_of_stream = 2; - // This field is used in ``GRPC`` body send mode when ``end_of_stream`` is - // true and ``body`` is empty. Those values would normally indicate an - // empty message on the stream with the end-of-stream bit set. - // However, if the half-close happens after the last message on the - // stream was already sent, then this field will be true to indicate an - // end-of-stream with *no* message (as opposed to an empty message). + // This field is used in ``GRPC`` body send mode when ``end_of_stream`` is ``true`` and ``body`` + // is empty. Those values would normally indicate an empty message on the stream with the + // end-of-stream bit set. However, if the half-close happens after the last message on the stream + // was already sent, then this field will be ``true`` to indicate an end-of-stream with *no* + // message (as opposed to an empty message). bool end_of_stream_without_message = 3; - // This field is used in ``GRPC`` body send mode to indicate whether - // the message is compressed. This will never be set to true by gRPC - // but may be set to true by a proxy like Envoy. + // This field is used in ``GRPC`` body send mode to indicate whether the message is compressed. + // This will never be set to ``true`` by gRPC but may be set to ``true`` by a proxy like Envoy. bool grpc_message_compressed = 4; } @@ -352,13 +465,14 @@ message TrailersResponse { HeaderMutation header_mutation = 1; } -// This message is sent by the external server to the data plane after ``HttpHeaders`` -// to initiate local response streaming. The server may follow up with multiple messages containing ``body_response``. -// The server must indicate end of stream by setting ``end_of_stream`` to ``true`` in the ``headers_response`` -// or ``body_response`` message or by sending a ``trailers_response`` message. +// This message is sent by the external server to the data plane after ``HttpHeaders`` to initiate +// local response streaming. The server may follow up with multiple messages containing +// ``body_response``. The server must indicate end of stream by setting ``end_of_stream`` to +// ``true`` in the ``headers_response`` or ``body_response`` message or by sending a +// ``trailers_response`` message. message StreamedImmediateResponse { oneof response { - // Response headers to be sent downstream. The ":status" header must be set. + // Response headers to be sent downstream. The ``:status`` header must be set. HttpHeaders headers_response = 1; // Response body to be sent downstream. @@ -384,7 +498,7 @@ message CommonResponse { // further messages for this request or response even if the processing // mode is configured to do so. // - // When used in response to a request_headers or response_headers message, + // When used in response to a ``request_headers`` or ``response_headers`` message, // this status makes it possible to either completely replace the body // while discarding the original body, or to add a body to a message that // formerly did not have one. @@ -401,23 +515,22 @@ message CommonResponse { ResponseStatus status = 1 [(validate.rules).enum = {defined_only: true}]; // Instructions on how to manipulate the headers. When responding to an - // HttpBody request, header mutations will only take effect if - // the current processing mode for the body is BUFFERED. + // ``HttpBody`` request, header mutations will only take effect if the current processing mode + // for the body is ``BUFFERED``. HeaderMutation header_mutation = 2; - // Replace the body of the last message sent to the remote server on this - // stream. If responding to an HttpBody request, simply replace or clear - // the body chunk that was sent with that request. Body mutations may take - // effect in response either to ``header`` or ``body`` messages. When it is - // in response to ``header`` messages, it only take effect if the + // Replace the body of the last message sent to the remote server on this stream. If responding + // to an ``HttpBody`` request, simply replace or clear the body chunk that was sent with that + // request. Body mutations may take effect in response either to ``header`` or ``body`` messages. + // When it is in response to ``header`` messages, it only takes effect if the // :ref:`status ` - // is set to CONTINUE_AND_REPLACE. + // is set to ``CONTINUE_AND_REPLACE``. BodyMutation body_mutation = 3; // [#not-implemented-hide:] - // Add new trailers to the message. This may be used when responding to either a - // HttpHeaders or HttpBody message, but only if this message is returned - // along with the CONTINUE_AND_REPLACE status. + // Add new trailers to the message. This may be used when responding to either an + // ``HttpHeaders`` or ``HttpBody`` message, but only if this message is returned + // along with the ``CONTINUE_AND_REPLACE`` status. // The header value is encoded in the // :ref:`raw_value ` field. config.core.v3.HeaderMap trailers = 4; @@ -429,34 +542,32 @@ message CommonResponse { bool clear_route_cache = 5; } -// This message causes the filter to attempt to create a locally -// generated response, send it downstream, stop processing -// additional filters, and ignore any additional messages received -// from the remote server for this request or response. If a response -// has already started, then this will either ship the reply directly -// to the downstream codec, or reset the stream. +// This message causes the filter to attempt to create a locally generated response, send it +// downstream, stop processing additional filters, and ignore any additional messages received +// from the remote server for this request or response. If a response has already started, then +// this will either ship the reply directly to the downstream codec, or reset the stream. // [#next-free-field: 6] message ImmediateResponse { // The response code to return. type.v3.HttpStatus status = 1 [(validate.rules).message = {required: true}]; - // Apply changes to the default headers, which will include content-type. + // Apply changes to the default headers, which will include ``content-type``. HeaderMutation headers = 2; // The message body to return with the response which is sent using the - // text/plain content type, or encoded in the grpc-message header. + // ``text/plain`` content type, or encoded in the ``grpc-message`` header. bytes body = 3; // If set, then include a gRPC status trailer. GrpcStatus grpc_status = 4; // A string detailing why this local reply was sent, which may be included - // in log and debug output (e.g. this populates the %RESPONSE_CODE_DETAILS% + // in log and debug output (e.g., this populates the ``%RESPONSE_CODE_DETAILS%`` // command operator field for use in access logging). string details = 5; } -// This message specifies a gRPC status for an ImmediateResponse message. +// This message specifies a gRPC status for an ``ImmediateResponse`` message. message GrpcStatus { // The actual gRPC status. uint32 status = 1; @@ -484,26 +595,24 @@ message StreamedBodyResponse { // a serialized gRPC message to be passed to the upstream/downstream by the data plane. bytes body = 1; - // The server sets this flag to true if it has received a body request with - // :ref:`end_of_stream ` set to true, - // and this is the last chunk of body responses. - // Note that in ``GRPC`` body send mode, this allows the ext_proc - // server to tell the data plane to send a half close after a client - // message, which will result in discarding any other messages sent by - // the client application. + // The server sets this flag to ``true`` if it has received a body request with + // :ref:`end_of_stream ` set to + // ``true``, and this is the last chunk of body responses. + // + // Note that in ``GRPC`` body send mode, this allows the ext_proc server to tell the data plane + // to send a half close after a client message, which will result in discarding any other + // messages sent by the client application. bool end_of_stream = 2; - // This field is used in ``GRPC`` body send mode when ``end_of_stream`` is - // true and ``body`` is empty. Those values would normally indicate an - // empty message on the stream with the end-of-stream bit set. - // However, if the half-close happens after the last message on the - // stream was already sent, then this field will be true to indicate an - // end-of-stream with *no* message (as opposed to an empty message). + // This field is used in ``GRPC`` body send mode when ``end_of_stream`` is ``true`` and ``body`` + // is empty. Those values would normally indicate an empty message on the stream with the + // end-of-stream bit set. However, if the half-close happens after the last message on the stream + // was already sent, then this field will be ``true`` to indicate an end-of-stream with *no* + // message (as opposed to an empty message). bool end_of_stream_without_message = 3; - // This field is used in ``GRPC`` body send mode to indicate whether - // the message is compressed. This will never be set to true by gRPC - // but may be set to true by a proxy like Envoy. + // This field is used in ``GRPC`` body send mode to indicate whether the message is compressed. + // This will never be set to ``true`` by gRPC but may be set to ``true`` by a proxy like Envoy. bool grpc_message_compressed = 4; } @@ -517,11 +626,10 @@ message BodyMutation { // is not set to ``FULL_DUPLEX_STREAMED`` or ``GRPC``. bytes body = 1; - // Clear the corresponding body chunk. - // Should only be used when the corresponding ``BodySendMode`` in the + // Clear the corresponding body chunk. Should only be used when the corresponding + // ``BodySendMode`` in the // :ref:`processing_mode ` // is not set to ``FULL_DUPLEX_STREAMED`` or ``GRPC``. - // Clear the corresponding body chunk. bool clear_body = 2; // Must be used when the corresponding ``BodySendMode`` in the