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